Use Double instead of Float for portable ChatOptions

This change updates the type of portable chat options from Float to
Double. Affected options include:
- frequencyPenalty
- presencePenalty
- temperature
- topP

The motivation for this change is to simplify coding. In Java, Float
values require an "f" suffix (e.g., 0.5f), while Double values don't
need any suffix. This makes Double easier to type and reduces
potential errors from forgetting the "f" suffix.

APIs, tests, and documentation have been updated to reflect this
change.

Fixes gh-712

Signed-off-by: Thomas Vitale <ThomasVitale@users.noreply.github.com>
This commit is contained in:
Thomas Vitale
2024-09-08 23:24:48 +02:00
committed by Mark Pollack
parent 40714c984d
commit 4b123a7516
146 changed files with 731 additions and 717 deletions

View File

@@ -82,7 +82,7 @@ public class AnthropicChatModel extends AbstractToolCallSupport implements ChatM
public static final Integer DEFAULT_MAX_TOKENS = 500; public static final Integer DEFAULT_MAX_TOKENS = 500;
public static final Float DEFAULT_TEMPERATURE = 0.8f; public static final Double DEFAULT_TEMPERATURE = 0.8;
/** /**
* The lower-level API for the Anthropic service. * The lower-level API for the Anthropic service.

View File

@@ -48,8 +48,8 @@ public class AnthropicChatOptions implements ChatOptions, FunctionCallingOptions
private @JsonProperty("max_tokens") Integer maxTokens; private @JsonProperty("max_tokens") Integer maxTokens;
private @JsonProperty("metadata") ChatCompletionRequest.Metadata metadata; private @JsonProperty("metadata") ChatCompletionRequest.Metadata metadata;
private @JsonProperty("stop_sequences") List<String> stopSequences; private @JsonProperty("stop_sequences") List<String> stopSequences;
private @JsonProperty("temperature") Float temperature; private @JsonProperty("temperature") Double temperature;
private @JsonProperty("top_p") Float topP; private @JsonProperty("top_p") Double topP;
private @JsonProperty("top_k") Integer topK; private @JsonProperty("top_k") Integer topK;
/** /**
@@ -112,12 +112,12 @@ public class AnthropicChatOptions implements ChatOptions, FunctionCallingOptions
return this; return this;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.temperature = temperature; this.options.temperature = temperature;
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.topP = topP; this.options.topP = topP;
return this; return this;
} }
@@ -186,20 +186,20 @@ public class AnthropicChatOptions implements ChatOptions, FunctionCallingOptions
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }
@@ -236,13 +236,13 @@ public class AnthropicChatOptions implements ChatOptions, FunctionCallingOptions
@Override @Override
@JsonIgnore @JsonIgnore
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return null; return null;
} }
@Override @Override
@JsonIgnore @JsonIgnore
public Float getPresencePenalty() { public Double getPresencePenalty() {
return null; return null;
} }

View File

@@ -241,19 +241,19 @@ public class AnthropicApi {
@JsonProperty("metadata") Metadata metadata, @JsonProperty("metadata") Metadata metadata,
@JsonProperty("stop_sequences") List<String> stopSequences, @JsonProperty("stop_sequences") List<String> stopSequences,
@JsonProperty("stream") Boolean stream, @JsonProperty("stream") Boolean stream,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("top_p") Float topP, @JsonProperty("top_p") Double topP,
@JsonProperty("top_k") Integer topK, @JsonProperty("top_k") Integer topK,
@JsonProperty("tools") List<Tool> tools) { @JsonProperty("tools") List<Tool> tools) {
// @formatter:on // @formatter:on
public ChatCompletionRequest(String model, List<AnthropicMessage> messages, String system, Integer maxTokens, public ChatCompletionRequest(String model, List<AnthropicMessage> messages, String system, Integer maxTokens,
Float temperature, Boolean stream) { Double temperature, Boolean stream) {
this(model, messages, system, maxTokens, null, null, stream, temperature, null, null, null); this(model, messages, system, maxTokens, null, null, stream, temperature, null, null, null);
} }
public ChatCompletionRequest(String model, List<AnthropicMessage> messages, String system, Integer maxTokens, public ChatCompletionRequest(String model, List<AnthropicMessage> messages, String system, Integer maxTokens,
List<String> stopSequences, Float temperature, Boolean stream) { List<String> stopSequences, Double temperature, Boolean stream) {
this(model, messages, system, maxTokens, null, stopSequences, stream, temperature, null, null, null); this(model, messages, system, maxTokens, null, stopSequences, stream, temperature, null, null, null);
} }
@@ -292,9 +292,9 @@ public class AnthropicApi {
private Boolean stream = false; private Boolean stream = false;
private Float temperature; private Double temperature;
private Float topP; private Double topP;
private Integer topK; private Integer topK;
@@ -357,12 +357,12 @@ public class AnthropicApi {
return this; return this;
} }
public ChatCompletionRequestBuilder withTemperature(Float temperature) { public ChatCompletionRequestBuilder withTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
return this; return this;
} }
public ChatCompletionRequestBuilder withTopP(Float topP) { public ChatCompletionRequestBuilder withTopP(Double topP) {
this.topP = topP; this.topP = topP;
return this; return this;
} }

View File

@@ -120,7 +120,7 @@ class AnthropicChatModelIT {
@Test @Test
void streamingWithTokenUsage() { void streamingWithTokenUsage() {
var promptOptions = AnthropicChatOptions.builder().withTemperature(0f).build(); var promptOptions = AnthropicChatOptions.builder().withTemperature(0.0).build();
var prompt = new Prompt("List two colors of the Polish flag. Be brief.", promptOptions); var prompt = new Prompt("List two colors of the Polish flag. Be brief.", promptOptions);
var streamingTokenUsage = this.chatModel.stream(prompt).blockLast().getMetadata().getUsage(); var streamingTokenUsage = this.chatModel.stream(prompt).blockLast().getMetadata().getUsage();

View File

@@ -71,9 +71,9 @@ public class AnthropicChatModelObservationIT {
.withModel(AnthropicApi.ChatModel.CLAUDE_3_HAIKU.getValue()) .withModel(AnthropicApi.ChatModel.CLAUDE_3_HAIKU.getValue())
.withMaxTokens(2048) .withMaxTokens(2048)
.withStopSequences(List.of("this-is-the-end")) .withStopSequences(List.of("this-is-the-end"))
.withTemperature(0.7f) .withTemperature(0.7)
.withTopK(1) .withTopK(1)
.withTopP(1f) .withTopP(1.0)
.build(); .build();
Prompt prompt = new Prompt("Why does a raven look like a desk?", options); Prompt prompt = new Prompt("Why does a raven look like a desk?", options);
@@ -93,9 +93,9 @@ public class AnthropicChatModelObservationIT {
.withModel(AnthropicApi.ChatModel.CLAUDE_3_HAIKU.getValue()) .withModel(AnthropicApi.ChatModel.CLAUDE_3_HAIKU.getValue())
.withMaxTokens(2048) .withMaxTokens(2048)
.withStopSequences(List.of("this-is-the-end")) .withStopSequences(List.of("this-is-the-end"))
.withTemperature(0.7f) .withTemperature(0.7)
.withTopK(1) .withTopK(1)
.withTopP(1f) .withTopP(1.0)
.build(); .build();
Prompt prompt = new Prompt("Why does a raven look like a desk?", options); Prompt prompt = new Prompt("Why does a raven look like a desk?", options);

View File

@@ -31,7 +31,7 @@ public class ChatCompletionRequestTests {
public void createRequestWithChatOptions() { public void createRequestWithChatOptions() {
var client = new AnthropicChatModel(new AnthropicApi("TEST"), var client = new AnthropicChatModel(new AnthropicApi("TEST"),
AnthropicChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6f).build()); AnthropicChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6).build());
var request = client.createRequest(new Prompt("Test message content"), false); var request = client.createRequest(new Prompt("Test message content"), false);
@@ -39,16 +39,16 @@ public class ChatCompletionRequestTests {
assertThat(request.stream()).isFalse(); assertThat(request.stream()).isFalse();
assertThat(request.model()).isEqualTo("DEFAULT_MODEL"); assertThat(request.model()).isEqualTo("DEFAULT_MODEL");
assertThat(request.temperature()).isEqualTo(66.6f); assertThat(request.temperature()).isEqualTo(66.6);
request = client.createRequest(new Prompt("Test message content", request = client.createRequest(new Prompt("Test message content",
AnthropicChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9f).build()), true); AnthropicChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9).build()), true);
assertThat(request.messages()).hasSize(1); assertThat(request.messages()).hasSize(1);
assertThat(request.stream()).isTrue(); assertThat(request.stream()).isTrue();
assertThat(request.model()).isEqualTo("PROMPT_MODEL"); assertThat(request.model()).isEqualTo("PROMPT_MODEL");
assertThat(request.temperature()).isEqualTo(99.9f); assertThat(request.temperature()).isEqualTo(99.9);
} }
} }

View File

@@ -45,7 +45,7 @@ public class AnthropicApiIT {
Role.USER); Role.USER);
ResponseEntity<ChatCompletionResponse> response = anthropicApi ResponseEntity<ChatCompletionResponse> response = anthropicApi
.chatCompletionEntity(new ChatCompletionRequest(AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue(), .chatCompletionEntity(new ChatCompletionRequest(AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue(),
List.of(chatCompletionMessage), null, 100, 0.8f, false)); List.of(chatCompletionMessage), null, 100, 0.8, false));
System.out.println(response); System.out.println(response);
assertThat(response).isNotNull(); assertThat(response).isNotNull();
@@ -58,9 +58,8 @@ public class AnthropicApiIT {
AnthropicMessage chatCompletionMessage = new AnthropicMessage(List.of(new ContentBlock("Tell me a Joke?")), AnthropicMessage chatCompletionMessage = new AnthropicMessage(List.of(new ContentBlock("Tell me a Joke?")),
Role.USER); Role.USER);
Flux<ChatCompletionResponse> response = anthropicApi Flux<ChatCompletionResponse> response = anthropicApi.chatCompletionStream(new ChatCompletionRequest(
.chatCompletionStream(new ChatCompletionRequest(AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue(), AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue(), List.of(chatCompletionMessage), null, 100, 0.8, true));
List.of(chatCompletionMessage), null, 100, 0.8f, true));
assertThat(response).isNotNull(); assertThat(response).isNotNull();

View File

@@ -107,8 +107,8 @@ public class AnthropicApiLegacyToolIT {
Role.USER); Role.USER);
ChatCompletionRequest chatCompletionRequest = new ChatCompletionRequest( ChatCompletionRequest chatCompletionRequest = new ChatCompletionRequest(
AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue(), List.of(chatCompletionMessage), systemPrompt, 500, AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue(), List.of(chatCompletionMessage), systemPrompt, 500, 0.8,
0.8f, false); false);
ResponseEntity<ChatCompletionResponse> chatCompletion = doCall(chatCompletionRequest); ResponseEntity<ChatCompletionResponse> chatCompletion = doCall(chatCompletionRequest);
@@ -147,7 +147,7 @@ public class AnthropicApiLegacyToolIT {
AnthropicMessage chatCompletionMessage2 = new AnthropicMessage(List.of(new ContentBlock(content)), Role.USER); AnthropicMessage chatCompletionMessage2 = new AnthropicMessage(List.of(new ContentBlock(content)), Role.USER);
return doCall(new ChatCompletionRequest(AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue(), return doCall(new ChatCompletionRequest(AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue(),
List.of(chatCompletionMessage2), null, 500, 0.8f, false)); List.of(chatCompletionMessage2), null, 500, 0.8, false));
} }
} }

View File

@@ -108,7 +108,7 @@ public class AnthropicApiToolIT {
.withModel(AnthropicApi.ChatModel.CLAUDE_3_OPUS) .withModel(AnthropicApi.ChatModel.CLAUDE_3_OPUS)
.withMessages(messageConversation) .withMessages(messageConversation)
.withMaxTokens(1500) .withMaxTokens(1500)
.withTemperature(0.8f) .withTemperature(0.8)
.withTools(tools) .withTools(tools)
.build(); .build();

View File

@@ -89,6 +89,7 @@ import reactor.core.publisher.Mono;
* @author Christian Tzolov * @author Christian Tzolov
* @author Grogdunn * @author Grogdunn
* @author Benoit Moussaud * @author Benoit Moussaud
* @author Thomas Vitale
* @author luocongqiu * @author luocongqiu
* @author timostark * @author timostark
* @see ChatModel * @see ChatModel
@@ -98,7 +99,7 @@ public class AzureOpenAiChatModel extends AbstractToolCallSupport implements Cha
private static final String DEFAULT_DEPLOYMENT_NAME = "gpt-4o"; private static final String DEFAULT_DEPLOYMENT_NAME = "gpt-4o";
private static final Float DEFAULT_TEMPERATURE = 0.7f; private static final Double DEFAULT_TEMPERATURE = 0.7;
/** /**
* The {@link OpenAIClient} used to interact with the Azure OpenAI service. * The {@link OpenAIClient} used to interact with the Azure OpenAI service.
@@ -422,22 +423,22 @@ public class AzureOpenAiChatModel extends AbstractToolCallSupport implements Cha
mergedAzureOptions.setTemperature(fromAzureOptions.getTemperature()); mergedAzureOptions.setTemperature(fromAzureOptions.getTemperature());
if (mergedAzureOptions.getTemperature() == null && toSpringAiOptions.getTemperature() != null) { if (mergedAzureOptions.getTemperature() == null && toSpringAiOptions.getTemperature() != null) {
mergedAzureOptions.setTemperature(toSpringAiOptions.getTemperature().doubleValue()); mergedAzureOptions.setTemperature(toSpringAiOptions.getTemperature());
} }
mergedAzureOptions.setTopP(fromAzureOptions.getTopP()); mergedAzureOptions.setTopP(fromAzureOptions.getTopP());
if (mergedAzureOptions.getTopP() == null && toSpringAiOptions.getTopP() != null) { if (mergedAzureOptions.getTopP() == null && toSpringAiOptions.getTopP() != null) {
mergedAzureOptions.setTopP(toSpringAiOptions.getTopP().doubleValue()); mergedAzureOptions.setTopP(toSpringAiOptions.getTopP());
} }
mergedAzureOptions.setFrequencyPenalty(fromAzureOptions.getFrequencyPenalty()); mergedAzureOptions.setFrequencyPenalty(fromAzureOptions.getFrequencyPenalty());
if (mergedAzureOptions.getFrequencyPenalty() == null && toSpringAiOptions.getFrequencyPenalty() != null) { if (mergedAzureOptions.getFrequencyPenalty() == null && toSpringAiOptions.getFrequencyPenalty() != null) {
mergedAzureOptions.setFrequencyPenalty(toSpringAiOptions.getFrequencyPenalty().doubleValue()); mergedAzureOptions.setFrequencyPenalty(toSpringAiOptions.getFrequencyPenalty());
} }
mergedAzureOptions.setPresencePenalty(fromAzureOptions.getPresencePenalty()); mergedAzureOptions.setPresencePenalty(fromAzureOptions.getPresencePenalty());
if (mergedAzureOptions.getPresencePenalty() == null && toSpringAiOptions.getPresencePenalty() != null) { if (mergedAzureOptions.getPresencePenalty() == null && toSpringAiOptions.getPresencePenalty() != null) {
mergedAzureOptions.setPresencePenalty(toSpringAiOptions.getPresencePenalty().doubleValue()); mergedAzureOptions.setPresencePenalty(toSpringAiOptions.getPresencePenalty());
} }
mergedAzureOptions.setResponseFormat(fromAzureOptions.getResponseFormat()); mergedAzureOptions.setResponseFormat(fromAzureOptions.getResponseFormat());
@@ -486,19 +487,19 @@ public class AzureOpenAiChatModel extends AbstractToolCallSupport implements Cha
} }
if (fromSpringAiOptions.getTemperature() != null) { if (fromSpringAiOptions.getTemperature() != null) {
mergedAzureOptions.setTemperature(fromSpringAiOptions.getTemperature().doubleValue()); mergedAzureOptions.setTemperature(fromSpringAiOptions.getTemperature());
} }
if (fromSpringAiOptions.getTopP() != null) { if (fromSpringAiOptions.getTopP() != null) {
mergedAzureOptions.setTopP(fromSpringAiOptions.getTopP().doubleValue()); mergedAzureOptions.setTopP(fromSpringAiOptions.getTopP());
} }
if (fromSpringAiOptions.getFrequencyPenalty() != null) { if (fromSpringAiOptions.getFrequencyPenalty() != null) {
mergedAzureOptions.setFrequencyPenalty(fromSpringAiOptions.getFrequencyPenalty().doubleValue()); mergedAzureOptions.setFrequencyPenalty(fromSpringAiOptions.getFrequencyPenalty());
} }
if (fromSpringAiOptions.getPresencePenalty() != null) { if (fromSpringAiOptions.getPresencePenalty() != null) {
mergedAzureOptions.setPresencePenalty(fromSpringAiOptions.getPresencePenalty().doubleValue()); mergedAzureOptions.setPresencePenalty(fromSpringAiOptions.getPresencePenalty());
} }
if (fromSpringAiOptions.getN() != null) { if (fromSpringAiOptions.getN() != null) {

View File

@@ -57,7 +57,7 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio
* two settings is difficult to predict. * two settings is difficult to predict.
*/ */
@JsonProperty(value = "temperature") @JsonProperty(value = "temperature")
private Float temperature; private Double temperature;
/** /**
* An alternative to sampling with temperature called nucleus sampling. This value * An alternative to sampling with temperature called nucleus sampling. This value
@@ -68,7 +68,7 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio
* two settings is difficult to predict. * two settings is difficult to predict.
*/ */
@JsonProperty(value = "top_p") @JsonProperty(value = "top_p")
private Float topP; private Double topP;
/** /**
* A map between GPT token IDs and bias scores that influences the probability of * A map between GPT token IDs and bias scores that influences the probability of
@@ -109,7 +109,7 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio
* output new topics. * output new topics.
*/ */
@JsonProperty(value = "presence_penalty") @JsonProperty(value = "presence_penalty")
private Float presencePenalty; private Double presencePenalty;
/** /**
* A value that influences the probability of generated tokens appearing based on * A value that influences the probability of generated tokens appearing based on
@@ -118,7 +118,7 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio
* model repeating the same statements verbatim. * model repeating the same statements verbatim.
*/ */
@JsonProperty(value = "frequency_penalty") @JsonProperty(value = "frequency_penalty")
private Float frequencyPenalty; private Double frequencyPenalty;
/** /**
* The deployment name as defined in Azure Open AI Studio when creating a deployment * The deployment name as defined in Azure Open AI Studio when creating a deployment
@@ -182,7 +182,7 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio
return this; return this;
} }
public Builder withFrequencyPenalty(Float frequencyPenalty) { public Builder withFrequencyPenalty(Double frequencyPenalty) {
this.options.frequencyPenalty = frequencyPenalty; this.options.frequencyPenalty = frequencyPenalty;
return this; return this;
} }
@@ -202,7 +202,7 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio
return this; return this;
} }
public Builder withPresencePenalty(Float presencePenalty) { public Builder withPresencePenalty(Double presencePenalty) {
this.options.presencePenalty = presencePenalty; this.options.presencePenalty = presencePenalty;
return this; return this;
} }
@@ -212,12 +212,12 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio
return this; return this;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.temperature = temperature; this.options.temperature = temperature;
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.topP = topP; this.options.topP = topP;
return this; return this;
} }
@@ -309,20 +309,20 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio
} }
@Override @Override
public Float getPresencePenalty() { public Double getPresencePenalty() {
return this.presencePenalty; return this.presencePenalty;
} }
public void setPresencePenalty(Float presencePenalty) { public void setPresencePenalty(Double presencePenalty) {
this.presencePenalty = presencePenalty; this.presencePenalty = presencePenalty;
} }
@Override @Override
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return this.frequencyPenalty; return this.frequencyPenalty;
} }
public void setFrequencyPenalty(Float frequencyPenalty) { public void setFrequencyPenalty(Double frequencyPenalty) {
this.frequencyPenalty = frequencyPenalty; this.frequencyPenalty = frequencyPenalty;
} }
@@ -346,20 +346,20 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }

View File

@@ -44,14 +44,14 @@ public class AzureChatCompletionsOptionsTests {
var defaultOptions = AzureOpenAiChatOptions.builder() var defaultOptions = AzureOpenAiChatOptions.builder()
.withDeploymentName("DEFAULT_MODEL") .withDeploymentName("DEFAULT_MODEL")
.withTemperature(66.6f) .withTemperature(66.6)
.withFrequencyPenalty(696.9f) .withFrequencyPenalty(696.9)
.withPresencePenalty(969.6f) .withPresencePenalty(969.6)
.withLogitBias(Map.of("foo", 1)) .withLogitBias(Map.of("foo", 1))
.withMaxTokens(969) .withMaxTokens(969)
.withN(69) .withN(69)
.withStop(List.of("foo", "bar")) .withStop(List.of("foo", "bar"))
.withTopP(0.69f) .withTopP(0.69)
.withUser("user") .withUser("user")
.withResponseFormat(AzureOpenAiResponseFormat.TEXT) .withResponseFormat(AzureOpenAiResponseFormat.TEXT)
.build(); .build();
@@ -63,27 +63,27 @@ public class AzureChatCompletionsOptionsTests {
assertThat(requestOptions.getMessages()).hasSize(1); assertThat(requestOptions.getMessages()).hasSize(1);
assertThat(requestOptions.getModel()).isEqualTo("DEFAULT_MODEL"); assertThat(requestOptions.getModel()).isEqualTo("DEFAULT_MODEL");
assertThat(requestOptions.getTemperature()).isEqualTo(66.6f); assertThat(requestOptions.getTemperature()).isEqualTo(66.6);
assertThat(requestOptions.getFrequencyPenalty()).isEqualTo(696.9f); assertThat(requestOptions.getFrequencyPenalty()).isEqualTo(696.9);
assertThat(requestOptions.getPresencePenalty()).isEqualTo(969.6f); assertThat(requestOptions.getPresencePenalty()).isEqualTo(969.6);
assertThat(requestOptions.getLogitBias()).isEqualTo(Map.of("foo", 1)); assertThat(requestOptions.getLogitBias()).isEqualTo(Map.of("foo", 1));
assertThat(requestOptions.getMaxTokens()).isEqualTo(969); assertThat(requestOptions.getMaxTokens()).isEqualTo(969);
assertThat(requestOptions.getN()).isEqualTo(69); assertThat(requestOptions.getN()).isEqualTo(69);
assertThat(requestOptions.getStop()).isEqualTo(List.of("foo", "bar")); assertThat(requestOptions.getStop()).isEqualTo(List.of("foo", "bar"));
assertThat(requestOptions.getTopP()).isEqualTo(0.69f); assertThat(requestOptions.getTopP()).isEqualTo(0.69);
assertThat(requestOptions.getUser()).isEqualTo("user"); assertThat(requestOptions.getUser()).isEqualTo("user");
assertThat(requestOptions.getResponseFormat()).isInstanceOf(ChatCompletionsTextResponseFormat.class); assertThat(requestOptions.getResponseFormat()).isInstanceOf(ChatCompletionsTextResponseFormat.class);
var runtimeOptions = AzureOpenAiChatOptions.builder() var runtimeOptions = AzureOpenAiChatOptions.builder()
.withDeploymentName("PROMPT_MODEL") .withDeploymentName("PROMPT_MODEL")
.withTemperature(99.9f) .withTemperature(99.9)
.withFrequencyPenalty(100f) .withFrequencyPenalty(100.0)
.withPresencePenalty(100f) .withPresencePenalty(100.0)
.withLogitBias(Map.of("foo", 2)) .withLogitBias(Map.of("foo", 2))
.withMaxTokens(100) .withMaxTokens(100)
.withN(100) .withN(100)
.withStop(List.of("foo", "bar")) .withStop(List.of("foo", "bar"))
.withTopP(0.111f) .withTopP(0.111)
.withUser("user2") .withUser("user2")
.withResponseFormat(AzureOpenAiResponseFormat.JSON) .withResponseFormat(AzureOpenAiResponseFormat.JSON)
.build(); .build();
@@ -93,30 +93,31 @@ public class AzureChatCompletionsOptionsTests {
assertThat(requestOptions.getMessages()).hasSize(1); assertThat(requestOptions.getMessages()).hasSize(1);
assertThat(requestOptions.getModel()).isEqualTo("PROMPT_MODEL"); assertThat(requestOptions.getModel()).isEqualTo("PROMPT_MODEL");
assertThat(requestOptions.getTemperature()).isEqualTo(99.9f); assertThat(requestOptions.getTemperature()).isEqualTo(99.9);
assertThat(requestOptions.getFrequencyPenalty()).isEqualTo(100f); assertThat(requestOptions.getFrequencyPenalty()).isEqualTo(100.0);
assertThat(requestOptions.getPresencePenalty()).isEqualTo(100f); assertThat(requestOptions.getPresencePenalty()).isEqualTo(100.0);
assertThat(requestOptions.getLogitBias()).isEqualTo(Map.of("foo", 2)); assertThat(requestOptions.getLogitBias()).isEqualTo(Map.of("foo", 2));
assertThat(requestOptions.getMaxTokens()).isEqualTo(100); assertThat(requestOptions.getMaxTokens()).isEqualTo(100);
assertThat(requestOptions.getN()).isEqualTo(100); assertThat(requestOptions.getN()).isEqualTo(100);
assertThat(requestOptions.getStop()).isEqualTo(List.of("foo", "bar")); assertThat(requestOptions.getStop()).isEqualTo(List.of("foo", "bar"));
assertThat(requestOptions.getTopP()).isEqualTo(0.111f); assertThat(requestOptions.getTopP()).isEqualTo(0.111);
assertThat(requestOptions.getUser()).isEqualTo("user2"); assertThat(requestOptions.getUser()).isEqualTo("user2");
assertThat(requestOptions.getResponseFormat()).isInstanceOf(ChatCompletionsJsonResponseFormat.class); assertThat(requestOptions.getResponseFormat()).isInstanceOf(ChatCompletionsJsonResponseFormat.class);
} }
private static Stream<Arguments> providePresencePenaltyAndFrequencyPenaltyTest() { private static Stream<Arguments> providePresencePenaltyAndFrequencyPenaltyTest() {
return Stream.of(Arguments.of(0.0f, 0.0f), Arguments.of(0.0f, 1.0f), Arguments.of(1.0f, 0.0f), return Stream.of(Arguments.of(0.0, 0.0), Arguments.of(0.0, 1.0), Arguments.of(1.0, 0.0), Arguments.of(1.0, 1.0),
Arguments.of(1.0f, 1.0f), Arguments.of(1.0f, null), Arguments.of(null, 1.0f), Arguments.of(null, null)); Arguments.of(1.0, null), Arguments.of(null, 1.0), Arguments.of(null, null));
} }
@ParameterizedTest @ParameterizedTest
@MethodSource("providePresencePenaltyAndFrequencyPenaltyTest") @MethodSource("providePresencePenaltyAndFrequencyPenaltyTest")
public void createChatOptionsWithPresencePenaltyAndFrequencyPenalty(Float presencePenalty, Float frequencyPenalty) { public void createChatOptionsWithPresencePenaltyAndFrequencyPenalty(Double presencePenalty,
Double frequencyPenalty) {
var options = AzureOpenAiChatOptions.builder() var options = AzureOpenAiChatOptions.builder()
.withMaxTokens(800) .withMaxTokens(800)
.withTemperature(0.7F) .withTemperature(0.7)
.withTopP(0.95F) .withTopP(0.95)
.withPresencePenalty(presencePenalty) .withPresencePenalty(presencePenalty)
.withFrequencyPenalty(frequencyPenalty) .withFrequencyPenalty(frequencyPenalty)
.build(); .build();
@@ -125,14 +126,14 @@ public class AzureChatCompletionsOptionsTests {
assertThat(options.getPresencePenalty()).isEqualTo(null); assertThat(options.getPresencePenalty()).isEqualTo(null);
} }
else { else {
assertThat(options.getPresencePenalty().floatValue()).isEqualTo(presencePenalty); assertThat(options.getPresencePenalty()).isEqualTo(presencePenalty);
} }
if (frequencyPenalty == null) { if (frequencyPenalty == null) {
assertThat(options.getFrequencyPenalty()).isEqualTo(null); assertThat(options.getFrequencyPenalty()).isEqualTo(null);
} }
else { else {
assertThat(options.getFrequencyPenalty().floatValue()).isEqualTo(frequencyPenalty); assertThat(options.getFrequencyPenalty()).isEqualTo(frequencyPenalty);
} }
} }

View File

@@ -39,7 +39,7 @@ public class AnthropicChatOptions implements ChatOptions {
* responses from the generative. This value specifies default to be used by the backend while making the call to * responses from the generative. This value specifies default to be used by the backend while making the call to
* the generative. * the generative.
*/ */
private @JsonProperty("temperature") Float temperature; private @JsonProperty("temperature") Double temperature;
/** /**
* Specify the maximum number of tokens to use in the generated response. Note that the models may stop before * Specify the maximum number of tokens to use in the generated response. Note that the models may stop before
@@ -57,7 +57,7 @@ public class AnthropicChatOptions implements ChatOptions {
* The maximum cumulative probability of tokens to consider when sampling. The generative uses combined Top-k and * The maximum cumulative probability of tokens to consider when sampling. The generative uses combined Top-k and
* nucleus sampling. Nucleus sampling considers the smallest set of tokens whose probability sum is at least topP. * nucleus sampling. Nucleus sampling considers the smallest set of tokens whose probability sum is at least topP.
*/ */
private @JsonProperty("top_p") Float topP; private @JsonProperty("top_p") Double topP;
/** /**
* Configure up to four sequences that the generative recognizes. After a stop sequence, the generative stops * Configure up to four sequences that the generative recognizes. After a stop sequence, the generative stops
@@ -79,7 +79,7 @@ public class AnthropicChatOptions implements ChatOptions {
private final AnthropicChatOptions options = new AnthropicChatOptions(); private final AnthropicChatOptions options = new AnthropicChatOptions();
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.setTemperature(temperature); this.options.setTemperature(temperature);
return this; return this;
} }
@@ -94,7 +94,7 @@ public class AnthropicChatOptions implements ChatOptions {
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.setTopP(topP); this.options.setTopP(topP);
return this; return this;
} }
@@ -116,11 +116,11 @@ public class AnthropicChatOptions implements ChatOptions {
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@@ -153,11 +153,11 @@ public class AnthropicChatOptions implements ChatOptions {
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }
@@ -186,13 +186,13 @@ public class AnthropicChatOptions implements ChatOptions {
@Override @Override
@JsonIgnore @JsonIgnore
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return null; return null;
} }
@Override @Override
@JsonIgnore @JsonIgnore
public Float getPresencePenalty() { public Double getPresencePenalty() {
return null; return null;
} }

View File

@@ -48,7 +48,7 @@ public class BedrockAnthropicChatModel implements ChatModel, StreamingChatModel
public BedrockAnthropicChatModel(AnthropicChatBedrockApi chatApi) { public BedrockAnthropicChatModel(AnthropicChatBedrockApi chatApi) {
this(chatApi, this(chatApi,
AnthropicChatOptions.builder() AnthropicChatOptions.builder()
.withTemperature(0.8f) .withTemperature(0.8)
.withMaxTokensToSample(500) .withMaxTokensToSample(500)
.withTopK(10) .withTopK(10)
.withAnthropicVersion(AnthropicChatBedrockApi.DEFAULT_ANTHROPIC_VERSION) .withAnthropicVersion(AnthropicChatBedrockApi.DEFAULT_ANTHROPIC_VERSION)

View File

@@ -34,6 +34,7 @@ import org.springframework.util.Assert;
/** /**
* @author Christian Tzolov * @author Christian Tzolov
* @author Thomas Vitale
* @author Wei Jiang * @author Wei Jiang
* @since 0.8.0 * @since 0.8.0
*/ */
@@ -139,10 +140,10 @@ public class AnthropicChatBedrockApi extends
@JsonInclude(Include.NON_NULL) @JsonInclude(Include.NON_NULL)
public record AnthropicChatRequest( public record AnthropicChatRequest(
@JsonProperty("prompt") String prompt, @JsonProperty("prompt") String prompt,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("max_tokens_to_sample") Integer maxTokensToSample, @JsonProperty("max_tokens_to_sample") Integer maxTokensToSample,
@JsonProperty("top_k") Integer topK, @JsonProperty("top_k") Integer topK,
@JsonProperty("top_p") Float topP, @JsonProperty("top_p") Double topP,
@JsonProperty("stop_sequences") List<String> stopSequences, @JsonProperty("stop_sequences") List<String> stopSequences,
@JsonProperty("anthropic_version") String anthropicVersion) { @JsonProperty("anthropic_version") String anthropicVersion) {
@@ -152,10 +153,10 @@ public class AnthropicChatBedrockApi extends
public static class Builder { public static class Builder {
private final String prompt; private final String prompt;
private Float temperature;// = 0.7f; private Double temperature;// = 0.7;
private Integer maxTokensToSample;// = 500; private Integer maxTokensToSample;// = 500;
private Integer topK;// = 10; private Integer topK;// = 10;
private Float topP; private Double topP;
private List<String> stopSequences; private List<String> stopSequences;
private String anthropicVersion; private String anthropicVersion;
@@ -163,7 +164,7 @@ public class AnthropicChatBedrockApi extends
this.prompt = prompt; this.prompt = prompt;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
return this; return this;
} }
@@ -178,7 +179,7 @@ public class AnthropicChatBedrockApi extends
return this; return this;
} }
public Builder withTopP(Float tpoP) { public Builder withTopP(Double tpoP) {
this.topP = tpoP; this.topP = tpoP;
return this; return this;
} }

View File

@@ -38,7 +38,7 @@ public class Anthropic3ChatOptions implements ChatOptions {
* responses from the generative. This value specifies default to be used by the backend while making the call to * responses from the generative. This value specifies default to be used by the backend while making the call to
* the generative. * the generative.
*/ */
private @JsonProperty("temperature") Float temperature; private @JsonProperty("temperature") Double temperature;
/** /**
* Specify the maximum number of tokens to use in the generated response. Note that the models may stop before * Specify the maximum number of tokens to use in the generated response. Note that the models may stop before
@@ -56,7 +56,7 @@ public class Anthropic3ChatOptions implements ChatOptions {
* The maximum cumulative probability of tokens to consider when sampling. The generative uses combined Top-k and * The maximum cumulative probability of tokens to consider when sampling. The generative uses combined Top-k and
* nucleus sampling. Nucleus sampling considers the smallest set of tokens whose probability sum is at least topP. * nucleus sampling. Nucleus sampling considers the smallest set of tokens whose probability sum is at least topP.
*/ */
private @JsonProperty("top_p") Float topP; private @JsonProperty("top_p") Double topP;
/** /**
* Configure up to four sequences that the generative recognizes. After a stop sequence, the generative stops * Configure up to four sequences that the generative recognizes. After a stop sequence, the generative stops
@@ -78,7 +78,7 @@ public class Anthropic3ChatOptions implements ChatOptions {
private final Anthropic3ChatOptions options = new Anthropic3ChatOptions(); private final Anthropic3ChatOptions options = new Anthropic3ChatOptions();
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.setTemperature(temperature); this.options.setTemperature(temperature);
return this; return this;
} }
@@ -93,7 +93,7 @@ public class Anthropic3ChatOptions implements ChatOptions {
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.setTopP(topP); this.options.setTopP(topP);
return this; return this;
} }
@@ -115,11 +115,11 @@ public class Anthropic3ChatOptions implements ChatOptions {
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@@ -142,11 +142,11 @@ public class Anthropic3ChatOptions implements ChatOptions {
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }
@@ -175,13 +175,13 @@ public class Anthropic3ChatOptions implements ChatOptions {
@Override @Override
@JsonIgnore @JsonIgnore
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return null; return null;
} }
@Override @Override
@JsonIgnore @JsonIgnore
public Float getPresencePenalty() { public Double getPresencePenalty() {
return null; return null;
} }

View File

@@ -62,7 +62,7 @@ public class BedrockAnthropic3ChatModel implements ChatModel, StreamingChatModel
public BedrockAnthropic3ChatModel(Anthropic3ChatBedrockApi chatApi) { public BedrockAnthropic3ChatModel(Anthropic3ChatBedrockApi chatApi) {
this(chatApi, this(chatApi,
Anthropic3ChatOptions.builder() Anthropic3ChatOptions.builder()
.withTemperature(0.8f) .withTemperature(0.8)
.withMaxTokens(500) .withMaxTokens(500)
.withTopK(10) .withTopK(10)
.withAnthropicVersion(Anthropic3ChatBedrockApi.DEFAULT_ANTHROPIC_VERSION) .withAnthropicVersion(Anthropic3ChatBedrockApi.DEFAULT_ANTHROPIC_VERSION)

View File

@@ -41,6 +41,7 @@ import java.util.List;
* *
* @author Ben Middleton * @author Ben Middleton
* @author Christian Tzolov * @author Christian Tzolov
* @author Thomas Vitale
* @author Wei Jiang * @author Wei Jiang
* @since 1.0.0 * @since 1.0.0
*/ */
@@ -146,10 +147,10 @@ public class Anthropic3ChatBedrockApi extends
public record AnthropicChatRequest( public record AnthropicChatRequest(
@JsonProperty("messages") List<ChatCompletionMessage> messages, @JsonProperty("messages") List<ChatCompletionMessage> messages,
@JsonProperty("system") String system, @JsonProperty("system") String system,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("max_tokens") Integer maxTokens, @JsonProperty("max_tokens") Integer maxTokens,
@JsonProperty("top_k") Integer topK, @JsonProperty("top_k") Integer topK,
@JsonProperty("top_p") Float topP, @JsonProperty("top_p") Double topP,
@JsonProperty("stop_sequences") List<String> stopSequences, @JsonProperty("stop_sequences") List<String> stopSequences,
@JsonProperty("anthropic_version") String anthropicVersion) { @JsonProperty("anthropic_version") String anthropicVersion) {
@@ -160,10 +161,10 @@ public class Anthropic3ChatBedrockApi extends
public static class Builder { public static class Builder {
private final List<ChatCompletionMessage> messages; private final List<ChatCompletionMessage> messages;
private String system; private String system;
private Float temperature;// = 0.7f; private Double temperature;// = 0.7;
private Integer maxTokens;// = 500; private Integer maxTokens;// = 500;
private Integer topK;// = 10; private Integer topK;// = 10;
private Float topP; private Double topP;
private List<String> stopSequences; private List<String> stopSequences;
private String anthropicVersion; private String anthropicVersion;
@@ -175,7 +176,7 @@ public class Anthropic3ChatBedrockApi extends
this.system = system; this.system = system;
return this; return this;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
return this; return this;
} }
@@ -190,7 +191,7 @@ public class Anthropic3ChatBedrockApi extends
return this; return this;
} }
public Builder withTopP(Float tpoP) { public Builder withTopP(Double tpoP) {
this.topP = tpoP; this.topP = tpoP;
return this; return this;
} }

View File

@@ -40,13 +40,13 @@ public class BedrockCohereChatOptions implements ChatOptions {
* (optional) Use a lower value to decrease randomness in the response. Defaults to * (optional) Use a lower value to decrease randomness in the response. Defaults to
* 0.7. * 0.7.
*/ */
@JsonProperty("temperature") Float temperature; @JsonProperty("temperature") Double temperature;
/** /**
* (optional) The maximum cumulative probability of tokens to consider when sampling. * (optional) The maximum cumulative probability of tokens to consider when sampling.
* The generative uses combined Top-k and nucleus sampling. Nucleus sampling considers * The generative uses combined Top-k and nucleus sampling. Nucleus sampling considers
* the smallest set of tokens whose probability sum is at least topP. * the smallest set of tokens whose probability sum is at least topP.
*/ */
@JsonProperty("p") Float topP; @JsonProperty("p") Double topP;
/** /**
* (optional) Specify the number of token choices the generative uses to generate the * (optional) Specify the number of token choices the generative uses to generate the
* next token. * next token.
@@ -89,12 +89,12 @@ public class BedrockCohereChatOptions implements ChatOptions {
private final BedrockCohereChatOptions options = new BedrockCohereChatOptions(); private final BedrockCohereChatOptions options = new BedrockCohereChatOptions();
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.setTemperature(temperature); this.options.setTemperature(temperature);
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.setTopP(topP); this.options.setTopP(topP);
return this; return this;
} }
@@ -141,20 +141,20 @@ public class BedrockCohereChatOptions implements ChatOptions {
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }
@@ -225,13 +225,13 @@ public class BedrockCohereChatOptions implements ChatOptions {
@Override @Override
@JsonIgnore @JsonIgnore
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return null; return null;
} }
@Override @Override
@JsonIgnore @JsonIgnore
public Float getPresencePenalty() { public Double getPresencePenalty() {
return null; return null;
} }

View File

@@ -38,6 +38,7 @@ import org.springframework.util.Assert;
* https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-cohere.html * https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-cohere.html
* *
* @author Christian Tzolov * @author Christian Tzolov
* @author Thomas Vitale
* @author Wei Jiang * @author Wei Jiang
* @since 0.8.0 * @since 0.8.0
*/ */
@@ -130,8 +131,8 @@ public class CohereChatBedrockApi extends
@JsonInclude(Include.NON_NULL) @JsonInclude(Include.NON_NULL)
public record CohereChatRequest( public record CohereChatRequest(
@JsonProperty("prompt") String prompt, @JsonProperty("prompt") String prompt,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("p") Float topP, @JsonProperty("p") Double topP,
@JsonProperty("k") Integer topK, @JsonProperty("k") Integer topK,
@JsonProperty("max_tokens") Integer maxTokens, @JsonProperty("max_tokens") Integer maxTokens,
@JsonProperty("stop_sequences") List<String> stopSequences, @JsonProperty("stop_sequences") List<String> stopSequences,
@@ -204,8 +205,8 @@ public class CohereChatBedrockApi extends
*/ */
public static class Builder { public static class Builder {
private final String prompt; private final String prompt;
private Float temperature; private Double temperature;
private Float topP; private Double topP;
private Integer topK; private Integer topK;
private Integer maxTokens; private Integer maxTokens;
private List<String> stopSequences; private List<String> stopSequences;
@@ -219,12 +220,12 @@ public class CohereChatBedrockApi extends
this.prompt = prompt; this.prompt = prompt;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.topP = topP; this.topP = topP;
return this; return this;
} }

View File

@@ -51,8 +51,8 @@ public class BedrockAi21Jurassic2ChatModel implements ChatModel {
public BedrockAi21Jurassic2ChatModel(Ai21Jurassic2ChatBedrockApi chatApi) { public BedrockAi21Jurassic2ChatModel(Ai21Jurassic2ChatBedrockApi chatApi) {
this(chatApi, this(chatApi,
BedrockAi21Jurassic2ChatOptions.builder() BedrockAi21Jurassic2ChatOptions.builder()
.withTemperature(0.8f) .withTemperature(0.8)
.withTopP(0.9f) .withTopP(0.9)
.withMaxTokens(100) .withMaxTokens(100)
.build()); .build());
} }

View File

@@ -61,13 +61,13 @@ public class BedrockAi21Jurassic2ChatOptions implements ChatOptions {
* Modifies the distribution from which tokens are sampled. * Modifies the distribution from which tokens are sampled.
*/ */
@JsonProperty("temperature") @JsonProperty("temperature")
private Float temperature; private Double temperature;
/** /**
* Sample tokens from the corresponding top percentile of probability mass. * Sample tokens from the corresponding top percentile of probability mass.
*/ */
@JsonProperty("topP") @JsonProperty("topP")
private Float topP; private Double topP;
/** /**
* Return the top-K (topKReturn) alternative tokens. * Return the top-K (topKReturn) alternative tokens.
@@ -171,7 +171,7 @@ public class BedrockAi21Jurassic2ChatOptions implements ChatOptions {
* @return The temperature. * @return The temperature.
*/ */
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return temperature; return temperature;
} }
@@ -179,7 +179,7 @@ public class BedrockAi21Jurassic2ChatOptions implements ChatOptions {
* Sets the temperature for modifying the token sampling distribution. * Sets the temperature for modifying the token sampling distribution.
* @param temperature The temperature. * @param temperature The temperature.
*/ */
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@@ -189,7 +189,7 @@ public class BedrockAi21Jurassic2ChatOptions implements ChatOptions {
* @return The topP parameter. * @return The topP parameter.
*/ */
@Override @Override
public Float getTopP() { public Double getTopP() {
return topP; return topP;
} }
@@ -198,7 +198,7 @@ public class BedrockAi21Jurassic2ChatOptions implements ChatOptions {
* mass. * mass.
* @param topP The topP parameter. * @param topP The topP parameter.
*/ */
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }
@@ -238,12 +238,12 @@ public class BedrockAi21Jurassic2ChatOptions implements ChatOptions {
@Override @Override
@JsonIgnore @JsonIgnore
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return getFrequencyPenaltyOptions() != null ? getFrequencyPenaltyOptions().scale() : null; return getFrequencyPenaltyOptions() != null ? getFrequencyPenaltyOptions().scale() : null;
} }
@JsonIgnore @JsonIgnore
public void setFrequencyPenalty(Float frequencyPenalty) { public void setFrequencyPenalty(Double frequencyPenalty) {
if (frequencyPenalty != null) { if (frequencyPenalty != null) {
setFrequencyPenaltyOptions(Penalty.builder().scale(frequencyPenalty).build()); setFrequencyPenaltyOptions(Penalty.builder().scale(frequencyPenalty).build());
} }
@@ -267,12 +267,12 @@ public class BedrockAi21Jurassic2ChatOptions implements ChatOptions {
@Override @Override
@JsonIgnore @JsonIgnore
public Float getPresencePenalty() { public Double getPresencePenalty() {
return getPresencePenaltyOptions() != null ? getPresencePenaltyOptions().scale() : null; return getPresencePenaltyOptions() != null ? getPresencePenaltyOptions().scale() : null;
} }
@JsonIgnore @JsonIgnore
public void setPresencePenalty(Float presencePenalty) { public void setPresencePenalty(Double presencePenalty) {
if (presencePenalty != null) { if (presencePenalty != null) {
setPresencePenaltyOptions(Penalty.builder().scale(presencePenalty).build()); setPresencePenaltyOptions(Penalty.builder().scale(presencePenalty).build());
} }
@@ -344,12 +344,12 @@ public class BedrockAi21Jurassic2ChatOptions implements ChatOptions {
return this; return this;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
request.setTemperature(temperature); request.setTemperature(temperature);
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
request.setTopP(topP); request.setTopP(topP);
return this; return this;
} }
@@ -389,7 +389,7 @@ public class BedrockAi21Jurassic2ChatOptions implements ChatOptions {
* Penalty object for frequency, presence, and count penalties. * Penalty object for frequency, presence, and count penalties.
*/ */
@JsonInclude(JsonInclude.Include.NON_NULL) @JsonInclude(JsonInclude.Include.NON_NULL)
public record Penalty(@JsonProperty("scale") Float scale, @JsonProperty("applyToNumbers") Boolean applyToNumbers, public record Penalty(@JsonProperty("scale") Double scale, @JsonProperty("applyToNumbers") Boolean applyToNumbers,
@JsonProperty("applyToPunctuations") Boolean applyToPunctuations, @JsonProperty("applyToPunctuations") Boolean applyToPunctuations,
@JsonProperty("applyToStopwords") Boolean applyToStopwords, @JsonProperty("applyToStopwords") Boolean applyToStopwords,
@JsonProperty("applyToWhitespaces") Boolean applyToWhitespaces, @JsonProperty("applyToWhitespaces") Boolean applyToWhitespaces,
@@ -401,7 +401,7 @@ public class BedrockAi21Jurassic2ChatOptions implements ChatOptions {
public static class Builder { public static class Builder {
private Float scale; private Double scale;
// can't keep it null due to modelOptionsUtils#mapToClass convert null to // can't keep it null due to modelOptionsUtils#mapToClass convert null to
// false // false
@@ -415,7 +415,7 @@ public class BedrockAi21Jurassic2ChatOptions implements ChatOptions {
private Boolean applyToEmojis = true; private Boolean applyToEmojis = true;
public Builder scale(Float scale) { public Builder scale(Double scale) {
this.scale = scale; this.scale = scale;
return this; return this;
} }

View File

@@ -37,6 +37,7 @@ import software.amazon.awssdk.regions.Region;
* https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-jurassic2.html * https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-jurassic2.html
* *
* @author Christian Tzolov * @author Christian Tzolov
* @author Thomas Vitale
* @author Wei Jiang * @author Wei Jiang
* @since 0.8.0 * @since 0.8.0
*/ */
@@ -132,8 +133,8 @@ public class Ai21Jurassic2ChatBedrockApi extends
@JsonInclude(Include.NON_NULL) @JsonInclude(Include.NON_NULL)
public record Ai21Jurassic2ChatRequest( public record Ai21Jurassic2ChatRequest(
@JsonProperty("prompt") String prompt, @JsonProperty("prompt") String prompt,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("topP") Float topP, @JsonProperty("topP") Double topP,
@JsonProperty("maxTokens") Integer maxTokens, @JsonProperty("maxTokens") Integer maxTokens,
@JsonProperty("stopSequences") List<String> stopSequences, @JsonProperty("stopSequences") List<String> stopSequences,
@JsonProperty("countPenalty") IntegerScalePenalty countPenalty, @JsonProperty("countPenalty") IntegerScalePenalty countPenalty,
@@ -198,8 +199,8 @@ public class Ai21Jurassic2ChatBedrockApi extends
} }
public static class Builder { public static class Builder {
private String prompt; private String prompt;
private Float temperature; private Double temperature;
private Float topP; private Double topP;
private Integer maxTokens; private Integer maxTokens;
private List<String> stopSequences; private List<String> stopSequences;
private IntegerScalePenalty countPenalty; private IntegerScalePenalty countPenalty;
@@ -210,12 +211,12 @@ public class Ai21Jurassic2ChatBedrockApi extends
this.prompt = prompt; this.prompt = prompt;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.topP = topP; this.topP = topP;
return this; return this;
} }

View File

@@ -49,8 +49,7 @@ public class BedrockLlamaChatModel implements ChatModel, StreamingChatModel {
private final BedrockLlamaChatOptions defaultOptions; private final BedrockLlamaChatOptions defaultOptions;
public BedrockLlamaChatModel(LlamaChatBedrockApi chatApi) { public BedrockLlamaChatModel(LlamaChatBedrockApi chatApi) {
this(chatApi, this(chatApi, BedrockLlamaChatOptions.builder().withTemperature(0.8).withTopP(0.9).withMaxGenLen(100).build());
BedrockLlamaChatOptions.builder().withTemperature(0.8f).withTopP(0.9f).withMaxGenLen(100).build());
} }
public BedrockLlamaChatModel(LlamaChatBedrockApi chatApi, BedrockLlamaChatOptions options) { public BedrockLlamaChatModel(LlamaChatBedrockApi chatApi, BedrockLlamaChatOptions options) {

View File

@@ -35,13 +35,13 @@ public class BedrockLlamaChatOptions implements ChatOptions {
* The temperature value controls the randomness of the generated text. Use a lower * The temperature value controls the randomness of the generated text. Use a lower
* value to decrease randomness in the response. * value to decrease randomness in the response.
*/ */
private @JsonProperty("temperature") Float temperature; private @JsonProperty("temperature") Double temperature;
/** /**
* The topP value controls the diversity of the generated text. Use a lower value to * The topP value controls the diversity of the generated text. Use a lower value to
* ignore less probable options. Set to 0 or 1.0 to disable. * ignore less probable options. Set to 0 or 1.0 to disable.
*/ */
private @JsonProperty("top_p") Float topP; private @JsonProperty("top_p") Double topP;
/** /**
* The maximum length of the generated text. * The maximum length of the generated text.
@@ -56,12 +56,12 @@ public class BedrockLlamaChatOptions implements ChatOptions {
private BedrockLlamaChatOptions options = new BedrockLlamaChatOptions(); private BedrockLlamaChatOptions options = new BedrockLlamaChatOptions();
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.setTemperature(temperature); this.options.setTemperature(temperature);
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.setTopP(topP); this.options.setTopP(topP);
return this; return this;
} }
@@ -78,20 +78,20 @@ public class BedrockLlamaChatOptions implements ChatOptions {
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }
@@ -122,13 +122,13 @@ public class BedrockLlamaChatOptions implements ChatOptions {
@Override @Override
@JsonIgnore @JsonIgnore
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return null; return null;
} }
@Override @Override
@JsonIgnore @JsonIgnore
public Float getPresencePenalty() { public Double getPresencePenalty() {
return null; return null;
} }

View File

@@ -36,6 +36,7 @@ import java.time.Duration;
* https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-meta.html * https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-meta.html
* *
* @author Christian Tzolov * @author Christian Tzolov
* @author Thomas Vitale
* @author Wei Jiang * @author Wei Jiang
* @since 0.8.0 * @since 0.8.0
*/ */
@@ -119,8 +120,8 @@ public class LlamaChatBedrockApi extends
@JsonInclude(Include.NON_NULL) @JsonInclude(Include.NON_NULL)
public record LlamaChatRequest( public record LlamaChatRequest(
@JsonProperty("prompt") String prompt, @JsonProperty("prompt") String prompt,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("top_p") Float topP, @JsonProperty("top_p") Double topP,
@JsonProperty("max_gen_len") Integer maxGenLen) { @JsonProperty("max_gen_len") Integer maxGenLen) {
/** /**
@@ -134,20 +135,20 @@ public class LlamaChatBedrockApi extends
public static class Builder { public static class Builder {
private String prompt; private String prompt;
private Float temperature; private Double temperature;
private Float topP; private Double topP;
private Integer maxGenLen; private Integer maxGenLen;
public Builder(String prompt) { public Builder(String prompt) {
this.prompt = prompt; this.prompt = prompt;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.topP = topP; this.topP = topP;
return this; return this;
} }

View File

@@ -46,7 +46,7 @@ public class BedrockTitanChatModel implements ChatModel, StreamingChatModel {
private final BedrockTitanChatOptions defaultOptions; private final BedrockTitanChatOptions defaultOptions;
public BedrockTitanChatModel(TitanChatBedrockApi chatApi) { public BedrockTitanChatModel(TitanChatBedrockApi chatApi) {
this(chatApi, BedrockTitanChatOptions.builder().withTemperature(0.8f).build()); this(chatApi, BedrockTitanChatOptions.builder().withTemperature(0.8).build());
} }
public BedrockTitanChatModel(TitanChatBedrockApi chatApi, BedrockTitanChatOptions defaultOptions) { public BedrockTitanChatModel(TitanChatBedrockApi chatApi, BedrockTitanChatOptions defaultOptions) {

View File

@@ -37,12 +37,12 @@ public class BedrockTitanChatOptions implements ChatOptions {
/** /**
* The temperature value controls the randomness of the generated text. * The temperature value controls the randomness of the generated text.
*/ */
private @JsonProperty("temperature") Float temperature; private @JsonProperty("temperature") Double temperature;
/** /**
* The topP value controls the diversity of the generated text. Use a lower value to ignore less probable options. * The topP value controls the diversity of the generated text. Use a lower value to ignore less probable options.
*/ */
private @JsonProperty("topP") Float topP; private @JsonProperty("topP") Double topP;
/** /**
* Maximum number of tokens to generate. * Maximum number of tokens to generate.
@@ -63,12 +63,12 @@ public class BedrockTitanChatOptions implements ChatOptions {
private BedrockTitanChatOptions options = new BedrockTitanChatOptions(); private BedrockTitanChatOptions options = new BedrockTitanChatOptions();
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.temperature = temperature; this.options.temperature = temperature;
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.topP = topP; this.options.topP = topP;
return this; return this;
} }
@@ -90,20 +90,20 @@ public class BedrockTitanChatOptions implements ChatOptions {
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return temperature; return temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return topP; return topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }
@@ -143,13 +143,13 @@ public class BedrockTitanChatOptions implements ChatOptions {
@Override @Override
@JsonIgnore @JsonIgnore
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return null; return null;
} }
@Override @Override
@JsonIgnore @JsonIgnore
public Float getPresencePenalty() { public Double getPresencePenalty() {
return null; return null;
} }

View File

@@ -40,6 +40,7 @@ import org.springframework.ai.model.ChatModelDescription;
* https://docs.aws.amazon.com/bedrock/latest/userguide/titan-text-models.html * https://docs.aws.amazon.com/bedrock/latest/userguide/titan-text-models.html
* *
* @author Christian Tzolov * @author Christian Tzolov
* @author Thomas Vitale
* @author Wei Jiang * @author Wei Jiang
* @since 0.8.0 * @since 0.8.0
*/ */
@@ -134,8 +135,8 @@ public class TitanChatBedrockApi extends
*/ */
@JsonInclude(Include.NON_NULL) @JsonInclude(Include.NON_NULL)
public record TextGenerationConfig( public record TextGenerationConfig(
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("topP") Float topP, @JsonProperty("topP") Double topP,
@JsonProperty("maxTokenCount") Integer maxTokenCount, @JsonProperty("maxTokenCount") Integer maxTokenCount,
@JsonProperty("stopSequences") List<String> stopSequences) { @JsonProperty("stopSequences") List<String> stopSequences) {
} }
@@ -151,8 +152,8 @@ public class TitanChatBedrockApi extends
public static class Builder { public static class Builder {
private final String inputText; private final String inputText;
private Float temperature; private Double temperature;
private Float topP; private Double topP;
private Integer maxTokenCount; private Integer maxTokenCount;
private List<String> stopSequences; private List<String> stopSequences;
@@ -160,12 +161,12 @@ public class TitanChatBedrockApi extends
this.inputText = inputText; this.inputText = inputText;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.topP = topP; this.topP = topP;
return this; return this;
} }

View File

@@ -40,9 +40,9 @@ public class BedrockAnthropicCreateRequestTests {
var client = new BedrockAnthropicChatModel(anthropicChatApi, var client = new BedrockAnthropicChatModel(anthropicChatApi,
AnthropicChatOptions.builder() AnthropicChatOptions.builder()
.withTemperature(66.6f) .withTemperature(66.6)
.withTopK(66) .withTopK(66)
.withTopP(0.66f) .withTopP(0.66)
.withMaxTokensToSample(666) .withMaxTokensToSample(666)
.withAnthropicVersion("X.Y.Z") .withAnthropicVersion("X.Y.Z")
.withStopSequences(List.of("stop1", "stop2")) .withStopSequences(List.of("stop1", "stop2"))
@@ -51,17 +51,17 @@ public class BedrockAnthropicCreateRequestTests {
var request = client.createRequest(new Prompt("Test message content")); var request = client.createRequest(new Prompt("Test message content"));
assertThat(request.prompt()).isNotEmpty(); assertThat(request.prompt()).isNotEmpty();
assertThat(request.temperature()).isEqualTo(66.6f); assertThat(request.temperature()).isEqualTo(66.6);
assertThat(request.topK()).isEqualTo(66); assertThat(request.topK()).isEqualTo(66);
assertThat(request.topP()).isEqualTo(0.66f); assertThat(request.topP()).isEqualTo(0.66);
assertThat(request.maxTokensToSample()).isEqualTo(666); assertThat(request.maxTokensToSample()).isEqualTo(666);
assertThat(request.anthropicVersion()).isEqualTo("X.Y.Z"); assertThat(request.anthropicVersion()).isEqualTo("X.Y.Z");
assertThat(request.stopSequences()).containsExactly("stop1", "stop2"); assertThat(request.stopSequences()).containsExactly("stop1", "stop2");
request = client.createRequest(new Prompt("Test message content", request = client.createRequest(new Prompt("Test message content",
AnthropicChatOptions.builder() AnthropicChatOptions.builder()
.withTemperature(99.9f) .withTemperature(99.9)
.withTopP(0.99f) .withTopP(0.99)
.withMaxTokensToSample(999) .withMaxTokensToSample(999)
.withAnthropicVersion("zzz") .withAnthropicVersion("zzz")
.withStopSequences(List.of("stop3", "stop4")) .withStopSequences(List.of("stop3", "stop4"))
@@ -70,9 +70,9 @@ public class BedrockAnthropicCreateRequestTests {
)); ));
assertThat(request.prompt()).isNotEmpty(); assertThat(request.prompt()).isNotEmpty();
assertThat(request.temperature()).isEqualTo(99.9f); assertThat(request.temperature()).isEqualTo(99.9);
assertThat(request.topK()).as("unchanged from the default options").isEqualTo(66); assertThat(request.topK()).as("unchanged from the default options").isEqualTo(66);
assertThat(request.topP()).isEqualTo(0.99f); assertThat(request.topP()).isEqualTo(0.99);
assertThat(request.maxTokensToSample()).isEqualTo(999); assertThat(request.maxTokensToSample()).isEqualTo(999);
assertThat(request.anthropicVersion()).isEqualTo("zzz"); assertThat(request.anthropicVersion()).isEqualTo("zzz");
assertThat(request.stopSequences()).containsExactly("stop3", "stop4"); assertThat(request.stopSequences()).containsExactly("stop3", "stop4");

View File

@@ -52,7 +52,7 @@ public class AnthropicChatBedrockApiIT {
AnthropicChatRequest request = AnthropicChatRequest AnthropicChatRequest request = AnthropicChatRequest
.builder(String.format(AnthropicChatBedrockApi.PROMPT_TEMPLATE, "Name 3 famous pirates")) .builder(String.format(AnthropicChatBedrockApi.PROMPT_TEMPLATE, "Name 3 famous pirates"))
.withTemperature(0.8f) .withTemperature(0.8)
.withMaxTokensToSample(300) .withMaxTokensToSample(300)
.withTopK(10) .withTopK(10)
.build(); .build();
@@ -75,7 +75,7 @@ public class AnthropicChatBedrockApiIT {
AnthropicChatRequest request = AnthropicChatRequest AnthropicChatRequest request = AnthropicChatRequest
.builder(String.format(AnthropicChatBedrockApi.PROMPT_TEMPLATE, "Name 3 famous pirates")) .builder(String.format(AnthropicChatBedrockApi.PROMPT_TEMPLATE, "Name 3 famous pirates"))
.withTemperature(0.8f) .withTemperature(0.8)
.withMaxTokensToSample(300) .withMaxTokensToSample(300)
.withTopK(10) .withTopK(10)
.withStopSequences(List.of("\n\nHuman:")) .withStopSequences(List.of("\n\nHuman:"))

View File

@@ -39,9 +39,9 @@ public class BedrockAnthropic3CreateRequestTests {
var client = new BedrockAnthropic3ChatModel(anthropicChatApi, var client = new BedrockAnthropic3ChatModel(anthropicChatApi,
Anthropic3ChatOptions.builder() Anthropic3ChatOptions.builder()
.withTemperature(66.6f) .withTemperature(66.6)
.withTopK(66) .withTopK(66)
.withTopP(0.66f) .withTopP(0.66)
.withMaxTokens(666) .withMaxTokens(666)
.withAnthropicVersion("X.Y.Z") .withAnthropicVersion("X.Y.Z")
.withStopSequences(List.of("stop1", "stop2")) .withStopSequences(List.of("stop1", "stop2"))
@@ -50,17 +50,17 @@ public class BedrockAnthropic3CreateRequestTests {
var request = client.createRequest(new Prompt("Test message content")); var request = client.createRequest(new Prompt("Test message content"));
assertThat(request.messages()).isNotEmpty(); assertThat(request.messages()).isNotEmpty();
assertThat(request.temperature()).isEqualTo(66.6f); assertThat(request.temperature()).isEqualTo(66.6);
assertThat(request.topK()).isEqualTo(66); assertThat(request.topK()).isEqualTo(66);
assertThat(request.topP()).isEqualTo(0.66f); assertThat(request.topP()).isEqualTo(0.66);
assertThat(request.maxTokens()).isEqualTo(666); assertThat(request.maxTokens()).isEqualTo(666);
assertThat(request.anthropicVersion()).isEqualTo("X.Y.Z"); assertThat(request.anthropicVersion()).isEqualTo("X.Y.Z");
assertThat(request.stopSequences()).containsExactly("stop1", "stop2"); assertThat(request.stopSequences()).containsExactly("stop1", "stop2");
request = client.createRequest(new Prompt("Test message content", request = client.createRequest(new Prompt("Test message content",
Anthropic3ChatOptions.builder() Anthropic3ChatOptions.builder()
.withTemperature(99.9f) .withTemperature(99.9)
.withTopP(0.99f) .withTopP(0.99)
.withMaxTokens(999) .withMaxTokens(999)
.withAnthropicVersion("zzz") .withAnthropicVersion("zzz")
.withStopSequences(List.of("stop3", "stop4")) .withStopSequences(List.of("stop3", "stop4"))
@@ -69,9 +69,9 @@ public class BedrockAnthropic3CreateRequestTests {
)); ));
assertThat(request.messages()).isNotEmpty(); assertThat(request.messages()).isNotEmpty();
assertThat(request.temperature()).isEqualTo(99.9f); assertThat(request.temperature()).isEqualTo(99.9);
assertThat(request.topK()).as("unchanged from the default options").isEqualTo(66); assertThat(request.topK()).as("unchanged from the default options").isEqualTo(66);
assertThat(request.topP()).isEqualTo(0.99f); assertThat(request.topP()).isEqualTo(0.99);
assertThat(request.maxTokens()).isEqualTo(999); assertThat(request.maxTokens()).isEqualTo(999);
assertThat(request.anthropicVersion()).isEqualTo("zzz"); assertThat(request.anthropicVersion()).isEqualTo("zzz");
assertThat(request.stopSequences()).containsExactly("stop3", "stop4"); assertThat(request.stopSequences()).containsExactly("stop3", "stop4");

View File

@@ -57,7 +57,7 @@ public class Anthropic3ChatBedrockApiIT {
MediaContent anthropicMessage = new MediaContent("Name 3 famous pirates"); MediaContent anthropicMessage = new MediaContent("Name 3 famous pirates");
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage(List.of(anthropicMessage), Role.USER); ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage(List.of(anthropicMessage), Role.USER);
AnthropicChatRequest request = AnthropicChatRequest.builder(List.of(chatCompletionMessage)) AnthropicChatRequest request = AnthropicChatRequest.builder(List.of(chatCompletionMessage))
.withTemperature(0.8f) .withTemperature(0.8)
.withMaxTokens(300) .withMaxTokens(300)
.withTopK(10) .withTopK(10)
.withAnthropicVersion(DEFAULT_ANTHROPIC_VERSION) .withAnthropicVersion(DEFAULT_ANTHROPIC_VERSION)
@@ -97,7 +97,7 @@ public class Anthropic3ChatBedrockApiIT {
AnthropicChatRequest request = AnthropicChatRequest AnthropicChatRequest request = AnthropicChatRequest
.builder(List.of(chatCompletionInitialMessage, chatCompletionAssistantMessage, .builder(List.of(chatCompletionInitialMessage, chatCompletionAssistantMessage,
chatCompletionFollowupMessage)) chatCompletionFollowupMessage))
.withTemperature(0.8f) .withTemperature(0.8)
.withMaxTokens(400) .withMaxTokens(400)
.withTopK(10) .withTopK(10)
.withAnthropicVersion(DEFAULT_ANTHROPIC_VERSION) .withAnthropicVersion(DEFAULT_ANTHROPIC_VERSION)
@@ -123,7 +123,7 @@ public class Anthropic3ChatBedrockApiIT {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage(List.of(anthropicMessage), Role.USER); ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage(List.of(anthropicMessage), Role.USER);
AnthropicChatRequest request = AnthropicChatRequest.builder(List.of(chatCompletionMessage)) AnthropicChatRequest request = AnthropicChatRequest.builder(List.of(chatCompletionMessage))
.withTemperature(0.8f) .withTemperature(0.8)
.withMaxTokens(300) .withMaxTokens(300)
.withTopK(10) .withTopK(10)
.withAnthropicVersion(DEFAULT_ANTHROPIC_VERSION) .withAnthropicVersion(DEFAULT_ANTHROPIC_VERSION)

View File

@@ -47,9 +47,9 @@ public class BedrockCohereChatCreateRequestTests {
var client = new BedrockCohereChatModel(chatApi, var client = new BedrockCohereChatModel(chatApi,
BedrockCohereChatOptions.builder() BedrockCohereChatOptions.builder()
.withTemperature(66.6f) .withTemperature(66.6)
.withTopK(66) .withTopK(66)
.withTopP(0.66f) .withTopP(0.66)
.withMaxTokens(678) .withMaxTokens(678)
.withStopSequences(List.of("stop1", "stop2")) .withStopSequences(List.of("stop1", "stop2"))
.withReturnLikelihoods(ReturnLikelihoods.ALL) .withReturnLikelihoods(ReturnLikelihoods.ALL)
@@ -63,9 +63,9 @@ public class BedrockCohereChatCreateRequestTests {
assertThat(request.prompt()).isNotEmpty(); assertThat(request.prompt()).isNotEmpty();
assertThat(request.stream()).isTrue(); assertThat(request.stream()).isTrue();
assertThat(request.temperature()).isEqualTo(66.6f); assertThat(request.temperature()).isEqualTo(66.6);
assertThat(request.topK()).isEqualTo(66); assertThat(request.topK()).isEqualTo(66);
assertThat(request.topP()).isEqualTo(0.66f); assertThat(request.topP()).isEqualTo(0.66);
assertThat(request.maxTokens()).isEqualTo(678); assertThat(request.maxTokens()).isEqualTo(678);
assertThat(request.stopSequences()).containsExactly("stop1", "stop2"); assertThat(request.stopSequences()).containsExactly("stop1", "stop2");
assertThat(request.returnLikelihoods()).isEqualTo(ReturnLikelihoods.ALL); assertThat(request.returnLikelihoods()).isEqualTo(ReturnLikelihoods.ALL);
@@ -75,9 +75,9 @@ public class BedrockCohereChatCreateRequestTests {
request = client.createRequest(new Prompt("Test message content", request = client.createRequest(new Prompt("Test message content",
BedrockCohereChatOptions.builder() BedrockCohereChatOptions.builder()
.withTemperature(99.9f) .withTemperature(99.9)
.withTopK(99) .withTopK(99)
.withTopP(0.99f) .withTopP(0.99)
.withMaxTokens(888) .withMaxTokens(888)
.withStopSequences(List.of("stop3", "stop4")) .withStopSequences(List.of("stop3", "stop4"))
.withReturnLikelihoods(ReturnLikelihoods.GENERATION) .withReturnLikelihoods(ReturnLikelihoods.GENERATION)
@@ -92,9 +92,9 @@ public class BedrockCohereChatCreateRequestTests {
assertThat(request.prompt()).isNotEmpty(); assertThat(request.prompt()).isNotEmpty();
assertThat(request.stream()).isFalse(); assertThat(request.stream()).isFalse();
assertThat(request.temperature()).isEqualTo(99.9f); assertThat(request.temperature()).isEqualTo(99.9);
assertThat(request.topK()).isEqualTo(99); assertThat(request.topK()).isEqualTo(99);
assertThat(request.topP()).isEqualTo(0.99f); assertThat(request.topP()).isEqualTo(0.99);
assertThat(request.maxTokens()).isEqualTo(888); assertThat(request.maxTokens()).isEqualTo(888);
assertThat(request.stopSequences()).containsExactly("stop3", "stop4"); assertThat(request.stopSequences()).containsExactly("stop3", "stop4");
assertThat(request.returnLikelihoods()).isEqualTo(ReturnLikelihoods.GENERATION); assertThat(request.returnLikelihoods()).isEqualTo(ReturnLikelihoods.GENERATION);

View File

@@ -49,13 +49,13 @@ public class CohereChatBedrockApiIT {
public void requestBuilder() { public void requestBuilder() {
CohereChatRequest request1 = new CohereChatRequest( CohereChatRequest request1 = new CohereChatRequest(
"What is the capital of Bulgaria and what is the size? What it the national anthem?", 0.5f, 0.9f, 15, "What is the capital of Bulgaria and what is the size? What it the national anthem?", 0.5, 0.9, 15, 40,
40, List.of("END"), CohereChatRequest.ReturnLikelihoods.ALL, false, 1, null, Truncate.NONE); List.of("END"), CohereChatRequest.ReturnLikelihoods.ALL, false, 1, null, Truncate.NONE);
var request2 = CohereChatRequest var request2 = CohereChatRequest
.builder("What is the capital of Bulgaria and what is the size? What it the national anthem?") .builder("What is the capital of Bulgaria and what is the size? What it the national anthem?")
.withTemperature(0.5f) .withTemperature(0.5)
.withTopP(0.9f) .withTopP(0.9)
.withTopK(15) .withTopK(15)
.withMaxTokens(40) .withMaxTokens(40)
.withStopSequences(List.of("END")) .withStopSequences(List.of("END"))
@@ -75,8 +75,8 @@ public class CohereChatBedrockApiIT {
var request = CohereChatRequest var request = CohereChatRequest
.builder("What is the capital of Bulgaria and what is the size? What it the national anthem?") .builder("What is the capital of Bulgaria and what is the size? What it the national anthem?")
.withStream(false) .withStream(false)
.withTemperature(0.5f) .withTemperature(0.5)
.withTopP(0.8f) .withTopP(0.8)
.withTopK(15) .withTopK(15)
.withMaxTokens(100) .withMaxTokens(100)
.withStopSequences(List.of("END")) .withStopSequences(List.of("END"))
@@ -100,8 +100,8 @@ public class CohereChatBedrockApiIT {
var request = CohereChatRequest var request = CohereChatRequest
.builder("What is the capital of Bulgaria and what is the size? What it the national anthem?") .builder("What is the capital of Bulgaria and what is the size? What it the national anthem?")
.withStream(true) .withStream(true)
.withTemperature(0.5f) .withTemperature(0.5)
.withTopP(0.8f) .withTopP(0.8)
.withTopK(15) .withTopK(15)
.withMaxTokens(100) .withMaxTokens(100)
.withStopSequences(List.of("END")) .withStopSequences(List.of("END"))

View File

@@ -157,9 +157,9 @@ class BedrockAi21Jurassic2ChatModelIT {
Ai21Jurassic2ChatBedrockApi jurassic2ChatBedrockApi) { Ai21Jurassic2ChatBedrockApi jurassic2ChatBedrockApi) {
return new BedrockAi21Jurassic2ChatModel(jurassic2ChatBedrockApi, return new BedrockAi21Jurassic2ChatModel(jurassic2ChatBedrockApi,
BedrockAi21Jurassic2ChatOptions.builder() BedrockAi21Jurassic2ChatOptions.builder()
.withTemperature(0.5f) .withTemperature(0.5)
.withMaxTokens(100) .withMaxTokens(100)
.withTopP(0.9f) .withTopP(0.9)
.build()); .build());
} }

View File

@@ -44,8 +44,8 @@ public class Ai21Jurassic2ChatBedrockApiIT {
@Test @Test
public void chatCompletion() { public void chatCompletion() {
Ai21Jurassic2ChatRequest request = new Ai21Jurassic2ChatRequest("Give me the names of 3 famous pirates?", 0.9f, Ai21Jurassic2ChatRequest request = new Ai21Jurassic2ChatRequest("Give me the names of 3 famous pirates?", 0.9,
0.9f, 100, null, // List.of("END"), 0.9, 100, null, // List.of("END"),
new Ai21Jurassic2ChatRequest.IntegerScalePenalty(1, true, true, true, true, true), new Ai21Jurassic2ChatRequest.IntegerScalePenalty(1, true, true, true, true, true),
new Ai21Jurassic2ChatRequest.FloatScalePenalty(0.5f, true, true, true, true, true), new Ai21Jurassic2ChatRequest.FloatScalePenalty(0.5f, true, true, true, true, true),
new Ai21Jurassic2ChatRequest.IntegerScalePenalty(1, true, true, true, true, true)); new Ai21Jurassic2ChatRequest.IntegerScalePenalty(1, true, true, true, true, true));

View File

@@ -208,7 +208,7 @@ class BedrockLlamaChatModelIT {
@Bean @Bean
public BedrockLlamaChatModel llamaChatModel(LlamaChatBedrockApi llamaApi) { public BedrockLlamaChatModel llamaChatModel(LlamaChatBedrockApi llamaApi) {
return new BedrockLlamaChatModel(llamaApi, return new BedrockLlamaChatModel(llamaApi,
BedrockLlamaChatOptions.builder().withTemperature(0.5f).withMaxGenLen(100).withTopP(0.9f).build()); BedrockLlamaChatOptions.builder().withTemperature(0.5).withMaxGenLen(100).withTopP(0.9).build());
} }
} }

View File

@@ -46,7 +46,7 @@ public class BedrockLlamaCreateRequestTests {
public void createRequestWithChatOptions() { public void createRequestWithChatOptions() {
var client = new BedrockLlamaChatModel(api, var client = new BedrockLlamaChatModel(api,
BedrockLlamaChatOptions.builder().withTemperature(66.6f).withMaxGenLen(666).withTopP(0.66f).build()); BedrockLlamaChatOptions.builder().withTemperature(66.6).withMaxGenLen(666).withTopP(0.66).build());
var request = client.createRequest(new Prompt("Test message content")); var request = client.createRequest(new Prompt("Test message content"));
@@ -56,7 +56,7 @@ public class BedrockLlamaCreateRequestTests {
assertThat(request.maxGenLen()).isEqualTo(666); assertThat(request.maxGenLen()).isEqualTo(666);
request = client.createRequest(new Prompt("Test message content", request = client.createRequest(new Prompt("Test message content",
BedrockLlamaChatOptions.builder().withTemperature(99.9f).withMaxGenLen(999).withTopP(0.99f).build())); BedrockLlamaChatOptions.builder().withTemperature(99.9).withMaxGenLen(999).withTopP(0.99).build()));
assertThat(request.prompt()).isNotEmpty(); assertThat(request.prompt()).isNotEmpty();
assertThat(request.temperature()).isEqualTo(99.9f); assertThat(request.temperature()).isEqualTo(99.9f);

View File

@@ -48,8 +48,8 @@ public class LlamaChatBedrockApiIT {
public void chatCompletion() { public void chatCompletion() {
LlamaChatRequest request = LlamaChatRequest.builder("Hello, my name is") LlamaChatRequest request = LlamaChatRequest.builder("Hello, my name is")
.withTemperature(0.9f) .withTemperature(0.9)
.withTopP(0.9f) .withTopP(0.9)
.withMaxGenLen(20) .withMaxGenLen(20)
.build(); .build();
@@ -67,7 +67,7 @@ public class LlamaChatBedrockApiIT {
@Test @Test
public void chatCompletionStream() { public void chatCompletionStream() {
LlamaChatRequest request = new LlamaChatRequest("Hello, my name is", 0.9f, 0.9f, 20); LlamaChatRequest request = new LlamaChatRequest("Hello, my name is", 0.9, 0.9, 20);
Flux<LlamaChatResponse> responseStream = llamaChatApi.chatCompletionStream(request); Flux<LlamaChatResponse> responseStream = llamaChatApi.chatCompletionStream(request);
List<LlamaChatResponse> responses = responseStream.collectList().block(); List<LlamaChatResponse> responses = responseStream.collectList().block();

View File

@@ -43,8 +43,8 @@ public class BedrockTitanChatModelCreateRequestTests {
var model = new BedrockTitanChatModel(api, var model = new BedrockTitanChatModel(api,
BedrockTitanChatOptions.builder() BedrockTitanChatOptions.builder()
.withTemperature(66.6f) .withTemperature(66.6)
.withTopP(0.66f) .withTopP(0.66)
.withMaxTokenCount(666) .withMaxTokenCount(666)
.withStopSequences(List.of("stop1", "stop2")) .withStopSequences(List.of("stop1", "stop2"))
.build()); .build());
@@ -52,15 +52,15 @@ public class BedrockTitanChatModelCreateRequestTests {
var request = model.createRequest(new Prompt("Test message content")); var request = model.createRequest(new Prompt("Test message content"));
assertThat(request.inputText()).isNotEmpty(); assertThat(request.inputText()).isNotEmpty();
assertThat(request.textGenerationConfig().temperature()).isEqualTo(66.6f); assertThat(request.textGenerationConfig().temperature()).isEqualTo(66.6);
assertThat(request.textGenerationConfig().topP()).isEqualTo(0.66f); assertThat(request.textGenerationConfig().topP()).isEqualTo(0.66);
assertThat(request.textGenerationConfig().maxTokenCount()).isEqualTo(666); assertThat(request.textGenerationConfig().maxTokenCount()).isEqualTo(666);
assertThat(request.textGenerationConfig().stopSequences()).containsExactly("stop1", "stop2"); assertThat(request.textGenerationConfig().stopSequences()).containsExactly("stop1", "stop2");
request = model.createRequest(new Prompt("Test message content", request = model.createRequest(new Prompt("Test message content",
BedrockTitanChatOptions.builder() BedrockTitanChatOptions.builder()
.withTemperature(99.9f) .withTemperature(99.9)
.withTopP(0.99f) .withTopP(0.99)
.withMaxTokenCount(999) .withMaxTokenCount(999)
.withStopSequences(List.of("stop3", "stop4")) .withStopSequences(List.of("stop3", "stop4"))
.build() .build()
@@ -68,8 +68,8 @@ public class BedrockTitanChatModelCreateRequestTests {
)); ));
assertThat(request.inputText()).isNotEmpty(); assertThat(request.inputText()).isNotEmpty();
assertThat(request.textGenerationConfig().temperature()).isEqualTo(99.9f); assertThat(request.textGenerationConfig().temperature()).isEqualTo(99.9);
assertThat(request.textGenerationConfig().topP()).isEqualTo(0.99f); assertThat(request.textGenerationConfig().topP()).isEqualTo(0.99);
assertThat(request.textGenerationConfig().maxTokenCount()).isEqualTo(999); assertThat(request.textGenerationConfig().maxTokenCount()).isEqualTo(999);
assertThat(request.textGenerationConfig().stopSequences()).containsExactly("stop3", "stop4"); assertThat(request.textGenerationConfig().stopSequences()).containsExactly("stop3", "stop4");
} }

View File

@@ -45,8 +45,8 @@ public class TitanChatBedrockApiIT {
Duration.ofMinutes(2)); Duration.ofMinutes(2));
TitanChatRequest titanChatRequest = TitanChatRequest.builder("Give me the names of 3 famous pirates?") TitanChatRequest titanChatRequest = TitanChatRequest.builder("Give me the names of 3 famous pirates?")
.withTemperature(0.5f) .withTemperature(0.5)
.withTopP(0.9f) .withTopP(0.9)
.withMaxTokenCount(100) .withMaxTokenCount(100)
.withStopSequences(List.of("|")) .withStopSequences(List.of("|"))
.build(); .build();

View File

@@ -99,7 +99,7 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
*/ */
public MiniMaxChatModel(MiniMaxApi miniMaxApi) { public MiniMaxChatModel(MiniMaxApi miniMaxApi) {
this(miniMaxApi, this(miniMaxApi,
MiniMaxChatOptions.builder().withModel(MiniMaxApi.DEFAULT_CHAT_MODEL).withTemperature(0.7f).build()); MiniMaxChatOptions.builder().withModel(MiniMaxApi.DEFAULT_CHAT_MODEL).withTemperature(0.7).build());
} }
/** /**

View File

@@ -54,7 +54,7 @@ public class MiniMaxChatOptions implements FunctionCallingOptions, ChatOptions {
* Number between -2.0 and 2.0. Positive values penalize new tokens based on their existing * Number between -2.0 and 2.0. Positive values penalize new tokens based on their existing
* frequency in the text so far, decreasing the model's likelihood to repeat the same line verbatim. * frequency in the text so far, decreasing the model's likelihood to repeat the same line verbatim.
*/ */
private @JsonProperty("frequency_penalty") Float frequencyPenalty; private @JsonProperty("frequency_penalty") Double frequencyPenalty;
/** /**
* The maximum number of tokens to generate in the chat completion. The total length of input * The maximum number of tokens to generate in the chat completion. The total length of input
* tokens and generated tokens is limited by the model's context length. * tokens and generated tokens is limited by the model's context length.
@@ -69,7 +69,7 @@ public class MiniMaxChatOptions implements FunctionCallingOptions, ChatOptions {
* Number between -2.0 and 2.0. Positive values penalize new tokens based on whether they * Number between -2.0 and 2.0. Positive values penalize new tokens based on whether they
* appear in the text so far, increasing the model's likelihood to talk about new topics. * appear in the text so far, increasing the model's likelihood to talk about new topics.
*/ */
private @JsonProperty("presence_penalty") Float presencePenalty; private @JsonProperty("presence_penalty") Double presencePenalty;
/** /**
* An object specifying the format that the model must output. Setting to { "type": * An object specifying the format that the model must output. Setting to { "type":
* "json_object" } enables JSON mode, which guarantees the message the model generates is valid JSON. * "json_object" } enables JSON mode, which guarantees the message the model generates is valid JSON.
@@ -92,13 +92,13 @@ public class MiniMaxChatOptions implements FunctionCallingOptions, ChatOptions {
* more random, while lower values like 0.2 will make it more focused and deterministic. We generally recommend * more random, while lower values like 0.2 will make it more focused and deterministic. We generally recommend
* altering this or top_p but not both. * altering this or top_p but not both.
*/ */
private @JsonProperty("temperature") Float temperature; private @JsonProperty("temperature") Double temperature;
/** /**
* An alternative to sampling with temperature, called nucleus sampling, where the model considers the * An alternative to sampling with temperature, called nucleus sampling, where the model considers the
* results of the tokens with top_p probability mass. So 0.1 means only the tokens comprising the top 10% * results of the tokens with top_p probability mass. So 0.1 means only the tokens comprising the top 10%
* probability mass are considered. We generally recommend altering this or temperature but not both. * probability mass are considered. We generally recommend altering this or temperature but not both.
*/ */
private @JsonProperty("top_p") Float topP; private @JsonProperty("top_p") Double topP;
/** /**
* Mask the text information in the output that is easy to involve privacy issues, * Mask the text information in the output that is easy to involve privacy issues,
* including but not limited to email, domain name, link, ID number, home address, etc. * including but not limited to email, domain name, link, ID number, home address, etc.
@@ -165,7 +165,7 @@ public class MiniMaxChatOptions implements FunctionCallingOptions, ChatOptions {
return this; return this;
} }
public Builder withFrequencyPenalty(Float frequencyPenalty) { public Builder withFrequencyPenalty(Double frequencyPenalty) {
this.options.frequencyPenalty = frequencyPenalty; this.options.frequencyPenalty = frequencyPenalty;
return this; return this;
} }
@@ -180,7 +180,7 @@ public class MiniMaxChatOptions implements FunctionCallingOptions, ChatOptions {
return this; return this;
} }
public Builder withPresencePenalty(Float presencePenalty) { public Builder withPresencePenalty(Double presencePenalty) {
this.options.presencePenalty = presencePenalty; this.options.presencePenalty = presencePenalty;
return this; return this;
} }
@@ -200,12 +200,12 @@ public class MiniMaxChatOptions implements FunctionCallingOptions, ChatOptions {
return this; return this;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.temperature = temperature; this.options.temperature = temperature;
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.topP = topP; this.options.topP = topP;
return this; return this;
} }
@@ -258,11 +258,11 @@ public class MiniMaxChatOptions implements FunctionCallingOptions, ChatOptions {
} }
@Override @Override
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return this.frequencyPenalty; return this.frequencyPenalty;
} }
public void setFrequencyPenalty(Float frequencyPenalty) { public void setFrequencyPenalty(Double frequencyPenalty) {
this.frequencyPenalty = frequencyPenalty; this.frequencyPenalty = frequencyPenalty;
} }
@@ -284,11 +284,11 @@ public class MiniMaxChatOptions implements FunctionCallingOptions, ChatOptions {
} }
@Override @Override
public Float getPresencePenalty() { public Double getPresencePenalty() {
return this.presencePenalty; return this.presencePenalty;
} }
public void setPresencePenalty(Float presencePenalty) { public void setPresencePenalty(Double presencePenalty) {
this.presencePenalty = presencePenalty; this.presencePenalty = presencePenalty;
} }
@@ -328,20 +328,20 @@ public class MiniMaxChatOptions implements FunctionCallingOptions, ChatOptions {
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }

View File

@@ -45,6 +45,7 @@ import java.util.function.Predicate;
* <a href="https://www.minimaxi.com/document/guides/Embeddings">MiniMax Embedding API</a>. * <a href="https://www.minimaxi.com/document/guides/Embeddings">MiniMax Embedding API</a>.
* *
* @author Geng Rong * @author Geng Rong
* @author Thomas Vitale
* @since 1.0.0 M1 * @since 1.0.0 M1
*/ */
public class MiniMaxApi { public class MiniMaxApi {
@@ -247,16 +248,16 @@ public class MiniMaxApi {
public record ChatCompletionRequest ( public record ChatCompletionRequest (
@JsonProperty("messages") List<ChatCompletionMessage> messages, @JsonProperty("messages") List<ChatCompletionMessage> messages,
@JsonProperty("model") String model, @JsonProperty("model") String model,
@JsonProperty("frequency_penalty") Float frequencyPenalty, @JsonProperty("frequency_penalty") Double frequencyPenalty,
@JsonProperty("max_tokens") Integer maxTokens, @JsonProperty("max_tokens") Integer maxTokens,
@JsonProperty("n") Integer n, @JsonProperty("n") Integer n,
@JsonProperty("presence_penalty") Float presencePenalty, @JsonProperty("presence_penalty") Double presencePenalty,
@JsonProperty("response_format") ResponseFormat responseFormat, @JsonProperty("response_format") ResponseFormat responseFormat,
@JsonProperty("seed") Integer seed, @JsonProperty("seed") Integer seed,
@JsonProperty("stop") List<String> stop, @JsonProperty("stop") List<String> stop,
@JsonProperty("stream") Boolean stream, @JsonProperty("stream") Boolean stream,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("top_p") Float topP, @JsonProperty("top_p") Double topP,
@JsonProperty("mask_sensitive_info") Boolean maskSensitiveInfo, @JsonProperty("mask_sensitive_info") Boolean maskSensitiveInfo,
@JsonProperty("tools") List<FunctionTool> tools, @JsonProperty("tools") List<FunctionTool> tools,
@JsonProperty("tool_choice") Object toolChoice) { @JsonProperty("tool_choice") Object toolChoice) {
@@ -268,7 +269,7 @@ public class MiniMaxApi {
* @param model ID of the model to use. * @param model ID of the model to use.
* @param temperature What sampling temperature to use, between 0 and 1. * @param temperature What sampling temperature to use, between 0 and 1.
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Float temperature) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Double temperature) {
this(messages, model, null, null, null, null, this(messages, model, null, null, null, null,
null, null, null, false, temperature, null,null, null, null, null, false, temperature, null,null,
null, null); null, null);
@@ -283,7 +284,7 @@ public class MiniMaxApi {
* @param stream If set, partial message deltas will be sent.Tokens will be sent as data-only server-sent events * @param stream If set, partial message deltas will be sent.Tokens will be sent as data-only server-sent events
* as they become available, with the stream terminated by a data: [DONE] message. * as they become available, with the stream terminated by a data: [DONE] message.
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Float temperature, boolean stream) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Double temperature, boolean stream) {
this(messages, model, null, null, null, null, this(messages, model, null, null, null, null,
null, null, null, stream, temperature, null,null, null, null, null, stream, temperature, null,null,
null, null); null, null);
@@ -301,7 +302,7 @@ public class MiniMaxApi {
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model,
List<FunctionTool> tools, Object toolChoice) { List<FunctionTool> tools, Object toolChoice) {
this(messages, model, null, null, null, null, this(messages, model, null, null, null, null,
null, null, null, false, 0.8f, null,null, null, null, null, false, 0.8, null,null,
tools, toolChoice); tools, toolChoice);
} }

View File

@@ -34,7 +34,7 @@ public class ChatCompletionRequestTests {
public void createRequestWithChatOptions() { public void createRequestWithChatOptions() {
var client = new MiniMaxChatModel(new MiniMaxApi("TEST"), var client = new MiniMaxChatModel(new MiniMaxApi("TEST"),
MiniMaxChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6f).build()); MiniMaxChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6).build());
var request = client.createRequest(new Prompt("Test message content"), false); var request = client.createRequest(new Prompt("Test message content"), false);
@@ -42,16 +42,16 @@ public class ChatCompletionRequestTests {
assertThat(request.stream()).isFalse(); assertThat(request.stream()).isFalse();
assertThat(request.model()).isEqualTo("DEFAULT_MODEL"); assertThat(request.model()).isEqualTo("DEFAULT_MODEL");
assertThat(request.temperature()).isEqualTo(66.6f); assertThat(request.temperature()).isEqualTo(66.6);
request = client.createRequest(new Prompt("Test message content", request = client.createRequest(new Prompt("Test message content",
MiniMaxChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9f).build()), true); MiniMaxChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9).build()), true);
assertThat(request.messages()).hasSize(1); assertThat(request.messages()).hasSize(1);
assertThat(request.stream()).isTrue(); assertThat(request.stream()).isTrue();
assertThat(request.model()).isEqualTo("PROMPT_MODEL"); assertThat(request.model()).isEqualTo("PROMPT_MODEL");
assertThat(request.temperature()).isEqualTo(99.9f); assertThat(request.temperature()).isEqualTo(99.9);
} }
@Test @Test

View File

@@ -43,7 +43,7 @@ public class MiniMaxApiIT {
void chatCompletionEntity() { void chatCompletionEntity() {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER); ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER);
ResponseEntity<ChatCompletion> response = miniMaxApi ResponseEntity<ChatCompletion> response = miniMaxApi
.chatCompletionEntity(new ChatCompletionRequest(List.of(chatCompletionMessage), "glm-4-air", 0.7f, false)); .chatCompletionEntity(new ChatCompletionRequest(List.of(chatCompletionMessage), "glm-4-air", 0.7, false));
assertThat(response).isNotNull(); assertThat(response).isNotNull();
assertThat(response.getBody()).isNotNull(); assertThat(response.getBody()).isNotNull();
@@ -53,7 +53,7 @@ public class MiniMaxApiIT {
void chatCompletionStream() { void chatCompletionStream() {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER); ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER);
Flux<ChatCompletionChunk> response = miniMaxApi Flux<ChatCompletionChunk> response = miniMaxApi
.chatCompletionStream(new ChatCompletionRequest(List.of(chatCompletionMessage), "glm-4-air", 0.7f, true)); .chatCompletionStream(new ChatCompletionRequest(List.of(chatCompletionMessage), "glm-4-air", 0.7, true));
assertThat(response).isNotNull(); assertThat(response).isNotNull();
assertThat(response.collectList().block()).isNotNull(); assertThat(response.collectList().block()).isNotNull();

View File

@@ -117,7 +117,7 @@ public class MiniMaxApiToolFunctionCallIT {
} }
var functionResponseRequest = new ChatCompletionRequest(messages, var functionResponseRequest = new ChatCompletionRequest(messages,
org.springframework.ai.minimax.api.MiniMaxApi.ChatModel.ABAB_6_5_Chat.getValue(), 0.5F); org.springframework.ai.minimax.api.MiniMaxApi.ChatModel.ABAB_6_5_Chat.getValue(), 0.5);
ResponseEntity<ChatCompletion> chatCompletion2 = miniMaxApi.chatCompletionEntity(functionResponseRequest); ResponseEntity<ChatCompletion> chatCompletion2 = miniMaxApi.chatCompletionEntity(functionResponseRequest);

View File

@@ -102,8 +102,8 @@ public class MistralAiChatModel extends AbstractToolCallSupport implements ChatM
public MistralAiChatModel(MistralAiApi mistralAiApi) { public MistralAiChatModel(MistralAiApi mistralAiApi) {
this(mistralAiApi, this(mistralAiApi,
MistralAiChatOptions.builder() MistralAiChatOptions.builder()
.withTemperature(0.7f) .withTemperature(0.7)
.withTopP(1f) .withTopP(1.0)
.withSafePrompt(false) .withSafePrompt(false)
.withModel(MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue()) .withModel(MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue())
.build()); .build());

View File

@@ -52,7 +52,7 @@ public class MistralAiChatOptions implements FunctionCallingOptions, ChatOptions
* make the output more random, while lower values like 0.2 will make it more focused * make the output more random, while lower values like 0.2 will make it more focused
* and deterministic. We generally recommend altering this or top_p but not both. * and deterministic. We generally recommend altering this or top_p but not both.
*/ */
private @JsonProperty("temperature") Float temperature; private @JsonProperty("temperature") Double temperature;
/** /**
* Nucleus sampling, where the model considers the results of the tokens with top_p * Nucleus sampling, where the model considers the results of the tokens with top_p
@@ -60,7 +60,7 @@ public class MistralAiChatOptions implements FunctionCallingOptions, ChatOptions
* mass are considered. We generally recommend altering this or temperature but not * mass are considered. We generally recommend altering this or temperature but not
* both. * both.
*/ */
private @JsonProperty("top_p") Float topP; private @JsonProperty("top_p") Double topP;
/** /**
* The maximum number of tokens to generate in the completion. The token count of your * The maximum number of tokens to generate in the completion. The token count of your
@@ -173,12 +173,12 @@ public class MistralAiChatOptions implements FunctionCallingOptions, ChatOptions
return this; return this;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.setTemperature(temperature); this.options.setTemperature(temperature);
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.setTopP(topP); this.options.setTopP(topP);
return this; return this;
} }
@@ -299,20 +299,20 @@ public class MistralAiChatOptions implements FunctionCallingOptions, ChatOptions
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }
@@ -340,13 +340,13 @@ public class MistralAiChatOptions implements FunctionCallingOptions, ChatOptions
@Override @Override
@JsonIgnore @JsonIgnore
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return null; return null;
} }
@Override @Override
@JsonIgnore @JsonIgnore
public Float getPresencePenalty() { public Double getPresencePenalty() {
return null; return null;
} }

View File

@@ -378,8 +378,8 @@ public class MistralAiApi {
@JsonProperty("messages") List<ChatCompletionMessage> messages, @JsonProperty("messages") List<ChatCompletionMessage> messages,
@JsonProperty("tools") List<FunctionTool> tools, @JsonProperty("tools") List<FunctionTool> tools,
@JsonProperty("tool_choice") ToolChoice toolChoice, @JsonProperty("tool_choice") ToolChoice toolChoice,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("top_p") Float topP, @JsonProperty("top_p") Double topP,
@JsonProperty("max_tokens") Integer maxTokens, @JsonProperty("max_tokens") Integer maxTokens,
@JsonProperty("stream") Boolean stream, @JsonProperty("stream") Boolean stream,
@JsonProperty("safe_prompt") Boolean safePrompt, @JsonProperty("safe_prompt") Boolean safePrompt,
@@ -396,7 +396,7 @@ public class MistralAiApi {
* @param model ID of the model to use. * @param model ID of the model to use.
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model) {
this(model, messages, null, null, 0.7f, 1f, null, false, false, null, null, null); this(model, messages, null, null, 0.7, 1.0, null, false, false, null, null, null);
} }
/** /**
@@ -409,9 +409,9 @@ public class MistralAiApi {
* @param stream Whether to stream back partial progress. If set, tokens will be * @param stream Whether to stream back partial progress. If set, tokens will be
* sent * sent
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Float temperature, public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Double temperature,
boolean stream) { boolean stream) {
this(model, messages, null, null, temperature, 1f, null, stream, false, null, null, null); this(model, messages, null, null, temperature, 1.0, null, stream, false, null, null, null);
} }
/** /**
@@ -423,8 +423,8 @@ public class MistralAiApi {
* @param temperature What sampling temperature to use, between 0.0 and 1.0. * @param temperature What sampling temperature to use, between 0.0 and 1.0.
* *
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Float temperature) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Double temperature) {
this(model, messages, null, null, temperature, 1f, null, false, false, null, null, null); this(model, messages, null, null, temperature, 1.0, null, false, false, null, null, null);
} }
/** /**
@@ -439,7 +439,7 @@ public class MistralAiApi {
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, List<FunctionTool> tools, public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, List<FunctionTool> tools,
ToolChoice toolChoice) { ToolChoice toolChoice) {
this(model, messages, tools, toolChoice, null, 1f, null, false, false, null, null, null); this(model, messages, tools, toolChoice, null, 1.0, null, false, false, null, null, null);
} }
/** /**
@@ -447,7 +447,7 @@ public class MistralAiApi {
* stream. * stream.
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, Boolean stream) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, Boolean stream) {
this(null, messages, null, null, 0.7f, 1f, null, stream, false, null, null, null); this(null, messages, null, null, 0.7, 1.0, null, stream, false, null, null, null);
} }
/** /**

View File

@@ -41,7 +41,7 @@ public class MistralAiChatCompletionRequestTest {
assertThat(request.messages()).hasSize(1); assertThat(request.messages()).hasSize(1);
assertThat(request.topP()).isEqualTo(1); assertThat(request.topP()).isEqualTo(1);
assertThat(request.temperature()).isEqualTo(0.7f); assertThat(request.temperature()).isEqualTo(0.7);
assertThat(request.safePrompt()).isFalse(); assertThat(request.safePrompt()).isFalse();
assertThat(request.maxTokens()).isNull(); assertThat(request.maxTokens()).isNull();
assertThat(request.stream()).isFalse(); assertThat(request.stream()).isFalse();
@@ -50,13 +50,13 @@ public class MistralAiChatCompletionRequestTest {
@Test @Test
void chatCompletionRequestWithOptionsTest() { void chatCompletionRequestWithOptionsTest() {
var options = MistralAiChatOptions.builder().withTemperature(0.5f).withTopP(0.8f).build(); var options = MistralAiChatOptions.builder().withTemperature(0.5).withTopP(0.8).build();
var request = chatModel.createRequest(new Prompt("test content", options), true); var request = chatModel.createRequest(new Prompt("test content", options), true);
assertThat(request.messages().size()).isEqualTo(1); assertThat(request.messages().size()).isEqualTo(1);
assertThat(request.topP()).isEqualTo(0.8f); assertThat(request.topP()).isEqualTo(0.8);
assertThat(request.temperature()).isEqualTo(0.5f); assertThat(request.temperature()).isEqualTo(0.5);
assertThat(request.stream()).isTrue(); assertThat(request.stream()).isTrue();
} }

View File

@@ -70,8 +70,8 @@ public class MistralAiChatModelObservationIT {
.withModel(MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue()) .withModel(MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue())
.withMaxTokens(2048) .withMaxTokens(2048)
.withStop(List.of("this-is-the-end")) .withStop(List.of("this-is-the-end"))
.withTemperature(0.7f) .withTemperature(0.7)
.withTopP(1f) .withTopP(1.0)
.build(); .build();
Prompt prompt = new Prompt("Why does a raven look like a desk?", options); Prompt prompt = new Prompt("Why does a raven look like a desk?", options);
@@ -91,8 +91,8 @@ public class MistralAiChatModelObservationIT {
.withModel(MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue()) .withModel(MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue())
.withMaxTokens(2048) .withMaxTokens(2048)
.withStop(List.of("this-is-the-end")) .withStop(List.of("this-is-the-end"))
.withTemperature(0.7f) .withTemperature(0.7)
.withTopP(1f) .withTopP(1.0)
.build(); .build();
Prompt prompt = new Prompt("Why does a raven look like a desk?", options); Prompt prompt = new Prompt("Why does a raven look like a desk?", options);

View File

@@ -96,8 +96,8 @@ public class MistralAiRetryTests {
chatModel = new MistralAiChatModel(mistralAiApi, chatModel = new MistralAiChatModel(mistralAiApi,
MistralAiChatOptions.builder() MistralAiChatOptions.builder()
.withTemperature(0.7f) .withTemperature(0.7)
.withTopP(1f) .withTopP(1.0)
.withSafePrompt(false) .withSafePrompt(false)
.withModel(MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue()) .withModel(MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue())
.build(), .build(),

View File

@@ -46,7 +46,7 @@ public class MistralAiApiIT {
void chatCompletionEntity() { void chatCompletionEntity() {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER); ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER);
ResponseEntity<ChatCompletion> response = mistralAiApi.chatCompletionEntity(new ChatCompletionRequest( ResponseEntity<ChatCompletion> response = mistralAiApi.chatCompletionEntity(new ChatCompletionRequest(
List.of(chatCompletionMessage), MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue(), 0.8f, false)); List.of(chatCompletionMessage), MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue(), 0.8, false));
assertThat(response).isNotNull(); assertThat(response).isNotNull();
assertThat(response.getBody()).isNotNull(); assertThat(response.getBody()).isNotNull();
@@ -63,7 +63,7 @@ public class MistralAiApiIT {
""", Role.SYSTEM); """, Role.SYSTEM);
ResponseEntity<ChatCompletion> response = mistralAiApi.chatCompletionEntity(new ChatCompletionRequest( ResponseEntity<ChatCompletion> response = mistralAiApi.chatCompletionEntity(new ChatCompletionRequest(
List.of(systemMessage, userMessage), MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue(), 0.8f, false)); List.of(systemMessage, userMessage), MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue(), 0.8, false));
assertThat(response).isNotNull(); assertThat(response).isNotNull();
assertThat(response.getBody()).isNotNull(); assertThat(response.getBody()).isNotNull();
@@ -73,7 +73,7 @@ public class MistralAiApiIT {
void chatCompletionStream() { void chatCompletionStream() {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER); ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER);
Flux<ChatCompletionChunk> response = mistralAiApi.chatCompletionStream(new ChatCompletionRequest( Flux<ChatCompletionChunk> response = mistralAiApi.chatCompletionStream(new ChatCompletionRequest(
List.of(chatCompletionMessage), MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue(), 0.8f, true)); List.of(chatCompletionMessage), MistralAiApi.ChatModel.OPEN_MISTRAL_7B.getValue(), 0.8, true));
assertThat(response).isNotNull(); assertThat(response).isNotNull();
assertThat(response.collectList().block()).isNotNull(); assertThat(response.collectList().block()).isNotNull();

View File

@@ -131,7 +131,7 @@ public class MistralAiApiToolFunctionCallIT {
} }
} }
var functionResponseRequest = new ChatCompletionRequest(messages, MISTRAL_AI_CHAT_MODEL, 0.8f); var functionResponseRequest = new ChatCompletionRequest(messages, MISTRAL_AI_CHAT_MODEL, 0.8);
ResponseEntity<ChatCompletion> chatCompletion2 = completionApi ResponseEntity<ChatCompletion> chatCompletion2 = completionApi
.chatCompletionEntity(functionResponseRequest); .chatCompletionEntity(functionResponseRequest);

View File

@@ -53,7 +53,7 @@ public class MoonshotChatOptions implements FunctionCallingOptions, ChatOptions
* make the output more random, while lower values like 0.2 will make it more focused * make the output more random, while lower values like 0.2 will make it more focused
* and deterministic. We generally recommend altering this or top_p but not both. * and deterministic. We generally recommend altering this or top_p but not both.
*/ */
private @JsonProperty("temperature") Float temperature; private @JsonProperty("temperature") Double temperature;
/** /**
* An alternative to sampling with temperature, called nucleus sampling, where the * An alternative to sampling with temperature, called nucleus sampling, where the
@@ -61,7 +61,7 @@ public class MoonshotChatOptions implements FunctionCallingOptions, ChatOptions
* only the tokens comprising the top 10% probability mass are considered. We * only the tokens comprising the top 10% probability mass are considered. We
* generally recommend altering this or temperature but not both. * generally recommend altering this or temperature but not both.
*/ */
private @JsonProperty("top_p") Float topP; private @JsonProperty("top_p") Double topP;
/** /**
* How many chat completion choices to generate for each input message. Note that you * How many chat completion choices to generate for each input message. Note that you
@@ -75,14 +75,14 @@ public class MoonshotChatOptions implements FunctionCallingOptions, ChatOptions
* they appear in the text so far, increasing the model's likelihood to talk about new * they appear in the text so far, increasing the model's likelihood to talk about new
* topics. * topics.
*/ */
private @JsonProperty("presence_penalty") Float presencePenalty; private @JsonProperty("presence_penalty") Double presencePenalty;
/** /**
* Number between -2.0 and 2.0. Positive values penalize new tokens based on their * Number between -2.0 and 2.0. Positive values penalize new tokens based on their
* existing frequency in the text so far, decreasing the model's likelihood to repeat * existing frequency in the text so far, decreasing the model's likelihood to repeat
* the same line verbatim. * the same line verbatim.
*/ */
private @JsonProperty("frequency_penalty") Float frequencyPenalty; private @JsonProperty("frequency_penalty") Double frequencyPenalty;
/** /**
* Up to 5 sequences where the API will stop generating further tokens. * Up to 5 sequences where the API will stop generating further tokens.
@@ -182,12 +182,12 @@ public class MoonshotChatOptions implements FunctionCallingOptions, ChatOptions
return this; return this;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.temperature = temperature; this.options.temperature = temperature;
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.topP = topP; this.options.topP = topP;
return this; return this;
} }
@@ -197,12 +197,12 @@ public class MoonshotChatOptions implements FunctionCallingOptions, ChatOptions
return this; return this;
} }
public Builder withPresencePenalty(Float presencePenalty) { public Builder withPresencePenalty(Double presencePenalty) {
this.options.presencePenalty = presencePenalty; this.options.presencePenalty = presencePenalty;
return this; return this;
} }
public Builder withFrequencyPenalty(Float frequencyPenalty) { public Builder withFrequencyPenalty(Double frequencyPenalty) {
this.options.frequencyPenalty = frequencyPenalty; this.options.frequencyPenalty = frequencyPenalty;
return this; return this;
} }
@@ -260,11 +260,11 @@ public class MoonshotChatOptions implements FunctionCallingOptions, ChatOptions
} }
@Override @Override
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return this.frequencyPenalty; return this.frequencyPenalty;
} }
public void setFrequencyPenalty(Float frequencyPenalty) { public void setFrequencyPenalty(Double frequencyPenalty) {
this.frequencyPenalty = frequencyPenalty; this.frequencyPenalty = frequencyPenalty;
} }
@@ -286,11 +286,11 @@ public class MoonshotChatOptions implements FunctionCallingOptions, ChatOptions
} }
@Override @Override
public Float getPresencePenalty() { public Double getPresencePenalty() {
return this.presencePenalty; return this.presencePenalty;
} }
public void setPresencePenalty(Float presencePenalty) { public void setPresencePenalty(Double presencePenalty) {
this.presencePenalty = presencePenalty; this.presencePenalty = presencePenalty;
} }
@@ -314,20 +314,20 @@ public class MoonshotChatOptions implements FunctionCallingOptions, ChatOptions
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }

View File

@@ -48,6 +48,7 @@ import static org.springframework.ai.moonshot.api.MoonshotConstants.DEFAULT_BASE
* </p> * </p>
* *
* @author Geng Rong * @author Geng Rong
* @author Thomas Vitale
*/ */
public class MoonshotApi { public class MoonshotApi {
@@ -155,11 +156,11 @@ public class MoonshotApi {
@JsonProperty("messages") List<ChatCompletionMessage> messages, @JsonProperty("messages") List<ChatCompletionMessage> messages,
@JsonProperty("model") String model, @JsonProperty("model") String model,
@JsonProperty("max_tokens") Integer maxTokens, @JsonProperty("max_tokens") Integer maxTokens,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("top_p") Float topP, @JsonProperty("top_p") Double topP,
@JsonProperty("n") Integer n, @JsonProperty("n") Integer n,
@JsonProperty("frequency_penalty") Float frequencyPenalty, @JsonProperty("frequency_penalty") Double frequencyPenalty,
@JsonProperty("presence_penalty") Float presencePenalty, @JsonProperty("presence_penalty") Double presencePenalty,
@JsonProperty("stop") List<String> stop, @JsonProperty("stop") List<String> stop,
@JsonProperty("stream") Boolean stream, @JsonProperty("stream") Boolean stream,
@JsonProperty("tools") List<FunctionTool> tools, @JsonProperty("tools") List<FunctionTool> tools,
@@ -174,7 +175,7 @@ public class MoonshotApi {
* @param model ID of the model to use. * @param model ID of the model to use.
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model) {
this(messages, model, null, 0.3f, 1f, null, null, null, null, false, null, null); this(messages, model, null, 0.3, 1.0, null, null, null, null, false, null, null);
} }
/** /**
@@ -187,9 +188,9 @@ public class MoonshotApi {
* @param stream Whether to stream back partial progress. If set, tokens will be * @param stream Whether to stream back partial progress. If set, tokens will be
* sent * sent
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Float temperature, public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Double temperature,
boolean stream) { boolean stream) {
this(messages, model, null, temperature, 1f, null, null, null, null, stream, null, null); this(messages, model, null, temperature, 1.0, null, null, null, null, stream, null, null);
} }
/** /**
@@ -200,8 +201,8 @@ public class MoonshotApi {
* @param model ID of the model to use. * @param model ID of the model to use.
* @param temperature What sampling temperature to use, between 0.0 and 1.0. * @param temperature What sampling temperature to use, between 0.0 and 1.0.
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Float temperature) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Double temperature) {
this(messages, model, null, temperature, 1f, null, null, null, null, false, null, null); this(messages, model, null, temperature, 1.0, null, null, null, null, false, null, null);
} }
/** /**
@@ -216,7 +217,7 @@ public class MoonshotApi {
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, List<FunctionTool> tools, public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, List<FunctionTool> tools,
Object toolChoice) { Object toolChoice) {
this(messages, model, null, null, 1f, null, null, null, null, false, tools, toolChoice); this(messages, model, null, null, 1.0, null, null, null, null, false, tools, toolChoice);
} }
/** /**
@@ -224,7 +225,7 @@ public class MoonshotApi {
* stream. * stream.
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, Boolean stream) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, Boolean stream) {
this(messages, DEFAULT_CHAT_MODEL, null, 0.7f, 1F, null, null, null, null, stream, null, null); this(messages, DEFAULT_CHAT_MODEL, null, 0.7, 1.0, null, null, null, null, stream, null, null);
} }
/** /**

View File

@@ -45,12 +45,12 @@ public class MoonshotChatCompletionRequestTest {
@Test @Test
void chatCompletionRequestWithOptionsTest() { void chatCompletionRequestWithOptionsTest() {
var options = MoonshotChatOptions.builder().withTemperature(0.5f).withTopP(0.8f).build(); var options = MoonshotChatOptions.builder().withTemperature(0.5).withTopP(0.8).build();
var request = chatModel.createRequest(new Prompt("test content", options), true); var request = chatModel.createRequest(new Prompt("test content", options), true);
assertThat(request.messages().size()).isEqualTo(1); assertThat(request.messages().size()).isEqualTo(1);
assertThat(request.topP()).isEqualTo(0.8f); assertThat(request.topP()).isEqualTo(0.8);
assertThat(request.temperature()).isEqualTo(0.5f); assertThat(request.temperature()).isEqualTo(0.5);
assertThat(request.stream()).isTrue(); assertThat(request.stream()).isTrue();
} }

View File

@@ -85,8 +85,8 @@ public class MoonshotRetryTests {
chatModel = new MoonshotChatModel(moonshotApi, chatModel = new MoonshotChatModel(moonshotApi,
MoonshotChatOptions.builder() MoonshotChatOptions.builder()
.withTemperature(0.7f) .withTemperature(0.7)
.withTopP(1f) .withTopP(1.0)
.withModel(MoonshotApi.ChatModel.MOONSHOT_V1_32K.getValue()) .withModel(MoonshotApi.ChatModel.MOONSHOT_V1_32K.getValue())
.build(), .build(),
null, retryTemplate); null, retryTemplate);

View File

@@ -41,7 +41,7 @@ public class MoonshotApiIT {
void chatCompletionEntity() { void chatCompletionEntity() {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER); ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER);
ResponseEntity<ChatCompletion> response = moonshotApi.chatCompletionEntity(new ChatCompletionRequest( ResponseEntity<ChatCompletion> response = moonshotApi.chatCompletionEntity(new ChatCompletionRequest(
List.of(chatCompletionMessage), MoonshotApi.ChatModel.MOONSHOT_V1_8K.getValue(), 0.8f, false)); List.of(chatCompletionMessage), MoonshotApi.ChatModel.MOONSHOT_V1_8K.getValue(), 0.8, false));
assertThat(response).isNotNull(); assertThat(response).isNotNull();
assertThat(response.getBody()).isNotNull(); assertThat(response.getBody()).isNotNull();
@@ -58,7 +58,7 @@ public class MoonshotApiIT {
""", Role.SYSTEM); """, Role.SYSTEM);
ResponseEntity<ChatCompletion> response = moonshotApi.chatCompletionEntity(new ChatCompletionRequest( ResponseEntity<ChatCompletion> response = moonshotApi.chatCompletionEntity(new ChatCompletionRequest(
List.of(systemMessage, userMessage), MoonshotApi.ChatModel.MOONSHOT_V1_8K.getValue(), 0.8f, false)); List.of(systemMessage, userMessage), MoonshotApi.ChatModel.MOONSHOT_V1_8K.getValue(), 0.8, false));
assertThat(response).isNotNull(); assertThat(response).isNotNull();
assertThat(response.getBody()).isNotNull(); assertThat(response.getBody()).isNotNull();
@@ -68,7 +68,7 @@ public class MoonshotApiIT {
void chatCompletionStream() { void chatCompletionStream() {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER); ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER);
Flux<ChatCompletionChunk> response = moonshotApi.chatCompletionStream(new ChatCompletionRequest( Flux<ChatCompletionChunk> response = moonshotApi.chatCompletionStream(new ChatCompletionRequest(
List.of(chatCompletionMessage), MoonshotApi.ChatModel.MOONSHOT_V1_8K.getValue(), 0.8f, true)); List.of(chatCompletionMessage), MoonshotApi.ChatModel.MOONSHOT_V1_8K.getValue(), 0.8, true));
assertThat(response).isNotNull(); assertThat(response).isNotNull();
assertThat(response.collectList().block()).isNotNull(); assertThat(response.collectList().block()).isNotNull();

View File

@@ -125,7 +125,7 @@ public class MoonshotApiToolFunctionCallIT {
} }
var functionResponseRequest = new ChatCompletionRequest(messages, var functionResponseRequest = new ChatCompletionRequest(messages,
MoonshotApi.ChatModel.MOONSHOT_V1_8K.getValue(), 0.5F); MoonshotApi.ChatModel.MOONSHOT_V1_8K.getValue(), 0.5);
ResponseEntity<ChatCompletion> chatCompletion2 = moonshotApi.chatCompletionEntity(functionResponseRequest); ResponseEntity<ChatCompletion> chatCompletion2 = moonshotApi.chatCompletionEntity(functionResponseRequest);

View File

@@ -171,7 +171,7 @@ public class OllamaOptions implements FunctionCallingOptions, ChatOptions, Embed
* more diverse text, while a lower value (e.g., 0.5) will generate more focused and * more diverse text, while a lower value (e.g., 0.5) will generate more focused and
* conservative text. (Default: 0.9) * conservative text. (Default: 0.9)
*/ */
@JsonProperty("top_p") private Float topP; @JsonProperty("top_p") private Double topP;
/** /**
* Tail free sampling is used to reduce the impact of less probable tokens * Tail free sampling is used to reduce the impact of less probable tokens
@@ -195,24 +195,24 @@ public class OllamaOptions implements FunctionCallingOptions, ChatOptions, Embed
* The temperature of the model. Increasing the temperature will * The temperature of the model. Increasing the temperature will
* make the model answer more creatively. (Default: 0.8) * make the model answer more creatively. (Default: 0.8)
*/ */
@JsonProperty("temperature") private Float temperature; @JsonProperty("temperature") private Double temperature;
/** /**
* Sets how strongly to penalize repetitions. A higher value * Sets how strongly to penalize repetitions. A higher value
* (e.g., 1.5) will penalize repetitions more strongly, while a lower value (e.g., * (e.g., 1.5) will penalize repetitions more strongly, while a lower value (e.g.,
* 0.9) will be more lenient. (Default: 1.1) * 0.9) will be more lenient. (Default: 1.1)
*/ */
@JsonProperty("repeat_penalty") private Float repeatPenalty; @JsonProperty("repeat_penalty") private Double repeatPenalty;
/** /**
* (Default: 0.0) * (Default: 0.0)
*/ */
@JsonProperty("presence_penalty") private Float presencePenalty; @JsonProperty("presence_penalty") private Double presencePenalty;
/** /**
* (Default: 0.0) * (Default: 0.0)
*/ */
@JsonProperty("frequency_penalty") private Float frequencyPenalty; @JsonProperty("frequency_penalty") private Double frequencyPenalty;
/** /**
* Enable Mirostat sampling for controlling perplexity. (default: 0, 0 * Enable Mirostat sampling for controlling perplexity. (default: 0, 0
@@ -414,7 +414,7 @@ public class OllamaOptions implements FunctionCallingOptions, ChatOptions, Embed
return this; return this;
} }
public OllamaOptions withTopP(Float topP) { public OllamaOptions withTopP(Double topP) {
this.topP = topP; this.topP = topP;
return this; return this;
} }
@@ -434,22 +434,22 @@ public class OllamaOptions implements FunctionCallingOptions, ChatOptions, Embed
return this; return this;
} }
public OllamaOptions withTemperature(Float temperature) { public OllamaOptions withTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
return this; return this;
} }
public OllamaOptions withRepeatPenalty(Float repeatPenalty) { public OllamaOptions withRepeatPenalty(Double repeatPenalty) {
this.repeatPenalty = repeatPenalty; this.repeatPenalty = repeatPenalty;
return this; return this;
} }
public OllamaOptions withPresencePenalty(Float presencePenalty) { public OllamaOptions withPresencePenalty(Double presencePenalty) {
this.presencePenalty = presencePenalty; this.presencePenalty = presencePenalty;
return this; return this;
} }
public OllamaOptions withFrequencyPenalty(Float frequencyPenalty) { public OllamaOptions withFrequencyPenalty(Double frequencyPenalty) {
this.frequencyPenalty = frequencyPenalty; this.frequencyPenalty = frequencyPenalty;
return this; return this;
} }
@@ -664,11 +664,11 @@ public class OllamaOptions implements FunctionCallingOptions, ChatOptions, Embed
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }
@@ -697,37 +697,37 @@ public class OllamaOptions implements FunctionCallingOptions, ChatOptions, Embed
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
public Float getRepeatPenalty() { public Double getRepeatPenalty() {
return this.repeatPenalty; return this.repeatPenalty;
} }
public void setRepeatPenalty(Float repeatPenalty) { public void setRepeatPenalty(Double repeatPenalty) {
this.repeatPenalty = repeatPenalty; this.repeatPenalty = repeatPenalty;
} }
@Override @Override
public Float getPresencePenalty() { public Double getPresencePenalty() {
return this.presencePenalty; return this.presencePenalty;
} }
public void setPresencePenalty(Float presencePenalty) { public void setPresencePenalty(Double presencePenalty) {
this.presencePenalty = presencePenalty; this.presencePenalty = presencePenalty;
} }
@Override @Override
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return this.frequencyPenalty; return this.frequencyPenalty;
} }
public void setFrequencyPenalty(Float frequencyPenalty) { public void setFrequencyPenalty(Double frequencyPenalty) {
this.frequencyPenalty = frequencyPenalty; this.frequencyPenalty = frequencyPenalty;
} }

View File

@@ -137,7 +137,7 @@ class OllamaChatModelFunctionCallingIT extends BaseOllamaIT {
@Bean @Bean
public OllamaChatModel ollamaChat(OllamaApi ollamaApi) { public OllamaChatModel ollamaChat(OllamaApi ollamaApi) {
return new OllamaChatModel(ollamaApi, OllamaOptions.create().withModel(MODEL).withTemperature(0.9f)); return new OllamaChatModel(ollamaApi, OllamaOptions.create().withModel(MODEL).withTemperature(0.9));
} }
} }

View File

@@ -86,7 +86,7 @@ class OllamaChatModelIT extends BaseOllamaIT {
UserMessage userMessage = new UserMessage("Tell me about 5 famous pirates from the Golden Age of Piracy."); UserMessage userMessage = new UserMessage("Tell me about 5 famous pirates from the Golden Age of Piracy.");
// portable/generic options // portable/generic options
var portableOptions = ChatOptionsBuilder.builder().withTemperature(0.7f).build(); var portableOptions = ChatOptionsBuilder.builder().withTemperature(0.7).build();
Prompt prompt = new Prompt(List.of(userMessage, systemMessage), portableOptions); Prompt prompt = new Prompt(List.of(userMessage, systemMessage), portableOptions);
@@ -240,7 +240,7 @@ class OllamaChatModelIT extends BaseOllamaIT {
@Bean @Bean
public OllamaChatModel ollamaChat(OllamaApi ollamaApi) { public OllamaChatModel ollamaChat(OllamaApi ollamaApi) {
return new OllamaChatModel(ollamaApi, OllamaOptions.create().withModel(MODEL).withTemperature(0.9f)); return new OllamaChatModel(ollamaApi, OllamaOptions.create().withModel(MODEL).withTemperature(0.9));
} }
} }

View File

@@ -91,7 +91,7 @@ class OllamaChatModelMultimodalIT extends BaseOllamaIT {
@Bean @Bean
public OllamaChatModel ollamaChat(OllamaApi ollamaApi) { public OllamaChatModel ollamaChat(OllamaApi ollamaApi) {
return new OllamaChatModel(ollamaApi, OllamaOptions.create().withModel(MODEL).withTemperature(0.9f)); return new OllamaChatModel(ollamaApi, OllamaOptions.create().withModel(MODEL).withTemperature(0.9));
} }
} }

View File

@@ -67,13 +67,13 @@ public class OllamaChatModelObservationIT extends BaseOllamaIT {
void observationForChatOperation() { void observationForChatOperation() {
var options = OllamaOptions.builder() var options = OllamaOptions.builder()
.withModel(OllamaModel.MISTRAL.getName()) .withModel(OllamaModel.MISTRAL.getName())
.withFrequencyPenalty(0f) .withFrequencyPenalty(0.0)
.withNumPredict(2048) .withNumPredict(2048)
.withPresencePenalty(0f) .withPresencePenalty(0.0)
.withStop(List.of("this-is-the-end")) .withStop(List.of("this-is-the-end"))
.withTemperature(0.7f) .withTemperature(0.7)
.withTopK(1) .withTopK(1)
.withTopP(1f) .withTopP(1.0)
.build(); .build();
Prompt prompt = new Prompt("Why does a raven look like a desk?", options); Prompt prompt = new Prompt("Why does a raven look like a desk?", options);
@@ -91,13 +91,13 @@ public class OllamaChatModelObservationIT extends BaseOllamaIT {
void observationForStreamingChatOperation() { void observationForStreamingChatOperation() {
var options = OllamaOptions.builder() var options = OllamaOptions.builder()
.withModel(OllamaModel.MISTRAL.getName()) .withModel(OllamaModel.MISTRAL.getName())
.withFrequencyPenalty(0f) .withFrequencyPenalty(0.0)
.withNumPredict(2048) .withNumPredict(2048)
.withPresencePenalty(0f) .withPresencePenalty(0.0)
.withStop(List.of("this-is-the-end")) .withStop(List.of("this-is-the-end"))
.withTemperature(0.7f) .withTemperature(0.7)
.withTopK(1) .withTopK(1)
.withTopP(1f) .withTopP(1.0)
.build(); .build();
Prompt prompt = new Prompt("Why does a raven look like a desk?", options); Prompt prompt = new Prompt("Why does a raven look like a desk?", options);

View File

@@ -30,7 +30,7 @@ import static org.assertj.core.api.Assertions.assertThat;
public class OllamaChatRequestTests { public class OllamaChatRequestTests {
OllamaChatModel chatModel = new OllamaChatModel(new OllamaApi(), OllamaChatModel chatModel = new OllamaChatModel(new OllamaApi(),
new OllamaOptions().withModel("MODEL_NAME").withTopK(99).withTemperature(66.6f).withNumGPU(1)); new OllamaOptions().withModel("MODEL_NAME").withTopK(99).withTemperature(66.6).withNumGPU(1));
@Test @Test
public void createRequestWithDefaultOptions() { public void createRequestWithDefaultOptions() {
@@ -51,7 +51,7 @@ public class OllamaChatRequestTests {
public void createRequestWithPromptOllamaOptions() { public void createRequestWithPromptOllamaOptions() {
// Runtime options should override the default options. // Runtime options should override the default options.
OllamaOptions promptOptions = new OllamaOptions().withTemperature(0.8f).withTopP(0.5f).withNumGPU(2); OllamaOptions promptOptions = new OllamaOptions().withTemperature(0.8).withTopP(0.5).withNumGPU(2);
var request = chatModel.ollamaChatRequest(new Prompt("Test message content", promptOptions), true); var request = chatModel.ollamaChatRequest(new Prompt("Test message content", promptOptions), true);
@@ -73,9 +73,9 @@ public class OllamaChatRequestTests {
// Ollama runtime options. // Ollama runtime options.
ChatOptions portablePromptOptions = ChatOptionsBuilder.builder() ChatOptions portablePromptOptions = ChatOptionsBuilder.builder()
.withTemperature(0.9f) .withTemperature(0.9)
.withTopK(100) .withTopK(100)
.withTopP(0.6f) .withTopP(0.6)
.build(); .build();
var request = chatModel.ollamaChatRequest(new Prompt("Test message content", portablePromptOptions), true); var request = chatModel.ollamaChatRequest(new Prompt("Test message content", portablePromptOptions), true);

View File

@@ -38,8 +38,6 @@ import java.util.stream.Collectors;
import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThat;
;
/** /**
* @author Christian Tzolov * @author Christian Tzolov
* @author Thomas Vitale * @author Thomas Vitale
@@ -97,7 +95,7 @@ public class OllamaApiIT extends BaseOllamaIT {
.withContent("What is the capital of Bulgaria and what is the size? " .withContent("What is the capital of Bulgaria and what is the size? "
+ "What it the national anthem?") + "What it the national anthem?")
.build())) .build()))
.withOptions(OllamaOptions.create().withTemperature(0.9f)) .withOptions(OllamaOptions.create().withTemperature(0.9))
.build(); .build();
ChatResponse response = ollamaApi.chat(request); ChatResponse response = ollamaApi.chat(request);
@@ -119,7 +117,7 @@ public class OllamaApiIT extends BaseOllamaIT {
.withMessages(List.of(Message.builder(Role.USER) .withMessages(List.of(Message.builder(Role.USER)
.withContent("What is the capital of Bulgaria and what is the size? " + "What it the national anthem?") .withContent("What is the capital of Bulgaria and what is the size? " + "What it the national anthem?")
.build())) .build()))
.withOptions(OllamaOptions.create().withTemperature(0.9f).toMap()) .withOptions(OllamaOptions.create().withTemperature(0.9).toMap())
.build(); .build();
Flux<ChatResponse> response = ollamaApi.streamingChat(request); Flux<ChatResponse> response = ollamaApi.streamingChat(request);

View File

@@ -28,7 +28,7 @@ public class OllamaModelOptionsTests {
@Test @Test
public void testOptions() { public void testOptions() {
var options = OllamaOptions.create().withTemperature(3.14f).withTopK(30).withStop(List.of("a", "b", "c")); var options = OllamaOptions.create().withTemperature(3.14).withTopK(30).withStop(List.of("a", "b", "c"));
var optionsMap = options.toMap(); var optionsMap = options.toMap();
System.out.println(optionsMap); System.out.println(optionsMap);

View File

@@ -134,7 +134,7 @@ public class OpenAiChatModel extends AbstractToolCallSupport implements ChatMode
*/ */
public OpenAiChatModel(OpenAiApi openAiApi) { public OpenAiChatModel(OpenAiApi openAiApi) {
this(openAiApi, this(openAiApi,
OpenAiChatOptions.builder().withModel(OpenAiApi.DEFAULT_CHAT_MODEL).withTemperature(0.7f).build()); OpenAiChatOptions.builder().withModel(OpenAiApi.DEFAULT_CHAT_MODEL).withTemperature(0.7).build());
} }
/** /**

View File

@@ -57,7 +57,7 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions {
* Number between -2.0 and 2.0. Positive values penalize new tokens based on their existing * Number between -2.0 and 2.0. Positive values penalize new tokens based on their existing
* frequency in the text so far, decreasing the model's likelihood to repeat the same line verbatim. * frequency in the text so far, decreasing the model's likelihood to repeat the same line verbatim.
*/ */
private @JsonProperty("frequency_penalty") Float frequencyPenalty; private @JsonProperty("frequency_penalty") Double frequencyPenalty;
/** /**
* Modify the likelihood of specified tokens appearing in the completion. Accepts a JSON object * Modify the likelihood of specified tokens appearing in the completion. Accepts a JSON object
* that maps tokens (specified by their token ID in the tokenizer) to an associated bias value from -100 to 100. * that maps tokens (specified by their token ID in the tokenizer) to an associated bias value from -100 to 100.
@@ -90,7 +90,7 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions {
* Number between -2.0 and 2.0. Positive values penalize new tokens based on whether they * Number between -2.0 and 2.0. Positive values penalize new tokens based on whether they
* appear in the text so far, increasing the model's likelihood to talk about new topics. * appear in the text so far, increasing the model's likelihood to talk about new topics.
*/ */
private @JsonProperty("presence_penalty") Float presencePenalty; private @JsonProperty("presence_penalty") Double presencePenalty;
/** /**
* An object specifying the format that the model must output. Setting to { "type": * An object specifying the format that the model must output. Setting to { "type":
* "json_object" } enables JSON mode, which guarantees the message the model generates is valid JSON. * "json_object" } enables JSON mode, which guarantees the message the model generates is valid JSON.
@@ -117,13 +117,13 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions {
* more random, while lower values like 0.2 will make it more focused and deterministic. We generally recommend * more random, while lower values like 0.2 will make it more focused and deterministic. We generally recommend
* altering this or top_p but not both. * altering this or top_p but not both.
*/ */
private @JsonProperty("temperature") Float temperature; private @JsonProperty("temperature") Double temperature;
/** /**
* An alternative to sampling with temperature, called nucleus sampling, where the model considers the * An alternative to sampling with temperature, called nucleus sampling, where the model considers the
* results of the tokens with top_p probability mass. So 0.1 means only the tokens comprising the top 10% * results of the tokens with top_p probability mass. So 0.1 means only the tokens comprising the top 10%
* probability mass are considered. We generally recommend altering this or temperature but not both. * probability mass are considered. We generally recommend altering this or temperature but not both.
*/ */
private @JsonProperty("top_p") Float topP; private @JsonProperty("top_p") Double topP;
/** /**
* A list of tools the model may call. Currently, only functions are supported as a tool. Use this to * A list of tools the model may call. Currently, only functions are supported as a tool. Use this to
* provide a list of functions the model may generate JSON inputs for. * provide a list of functions the model may generate JSON inputs for.
@@ -205,7 +205,7 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions {
return this; return this;
} }
public Builder withFrequencyPenalty(Float frequencyPenalty) { public Builder withFrequencyPenalty(Double frequencyPenalty) {
this.options.frequencyPenalty = frequencyPenalty; this.options.frequencyPenalty = frequencyPenalty;
return this; return this;
} }
@@ -235,7 +235,7 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions {
return this; return this;
} }
public Builder withPresencePenalty(Float presencePenalty) { public Builder withPresencePenalty(Double presencePenalty) {
this.options.presencePenalty = presencePenalty; this.options.presencePenalty = presencePenalty;
return this; return this;
} }
@@ -260,12 +260,12 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions {
return this; return this;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.temperature = temperature; this.options.temperature = temperature;
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.topP = topP; this.options.topP = topP;
return this; return this;
} }
@@ -337,11 +337,11 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions {
} }
@Override @Override
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return this.frequencyPenalty; return this.frequencyPenalty;
} }
public void setFrequencyPenalty(Float frequencyPenalty) { public void setFrequencyPenalty(Double frequencyPenalty) {
this.frequencyPenalty = frequencyPenalty; this.frequencyPenalty = frequencyPenalty;
} }
@@ -387,11 +387,11 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions {
} }
@Override @Override
public Float getPresencePenalty() { public Double getPresencePenalty() {
return this.presencePenalty; return this.presencePenalty;
} }
public void setPresencePenalty(Float presencePenalty) { public void setPresencePenalty(Double presencePenalty) {
this.presencePenalty = presencePenalty; this.presencePenalty = presencePenalty;
} }
@@ -439,20 +439,20 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions {
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }

View File

@@ -415,20 +415,20 @@ public class OpenAiApi {
public record ChatCompletionRequest(// @formatter:off public record ChatCompletionRequest(// @formatter:off
@JsonProperty("messages") List<ChatCompletionMessage> messages, @JsonProperty("messages") List<ChatCompletionMessage> messages,
@JsonProperty("model") String model, @JsonProperty("model") String model,
@JsonProperty("frequency_penalty") Float frequencyPenalty, @JsonProperty("frequency_penalty") Double frequencyPenalty,
@JsonProperty("logit_bias") Map<String, Integer> logitBias, @JsonProperty("logit_bias") Map<String, Integer> logitBias,
@JsonProperty("logprobs") Boolean logprobs, @JsonProperty("logprobs") Boolean logprobs,
@JsonProperty("top_logprobs") Integer topLogprobs, @JsonProperty("top_logprobs") Integer topLogprobs,
@JsonProperty("max_tokens") Integer maxTokens, @JsonProperty("max_tokens") Integer maxTokens,
@JsonProperty("n") Integer n, @JsonProperty("n") Integer n,
@JsonProperty("presence_penalty") Float presencePenalty, @JsonProperty("presence_penalty") Double presencePenalty,
@JsonProperty("response_format") ResponseFormat responseFormat, @JsonProperty("response_format") ResponseFormat responseFormat,
@JsonProperty("seed") Integer seed, @JsonProperty("seed") Integer seed,
@JsonProperty("stop") List<String> stop, @JsonProperty("stop") List<String> stop,
@JsonProperty("stream") Boolean stream, @JsonProperty("stream") Boolean stream,
@JsonProperty("stream_options") StreamOptions streamOptions, @JsonProperty("stream_options") StreamOptions streamOptions,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("top_p") Float topP, @JsonProperty("top_p") Double topP,
@JsonProperty("tools") List<FunctionTool> tools, @JsonProperty("tools") List<FunctionTool> tools,
@JsonProperty("tool_choice") Object toolChoice, @JsonProperty("tool_choice") Object toolChoice,
@JsonProperty("parallel_tool_calls") Boolean parallelToolCalls, @JsonProperty("parallel_tool_calls") Boolean parallelToolCalls,
@@ -441,7 +441,7 @@ public class OpenAiApi {
* @param model ID of the model to use. * @param model ID of the model to use.
* @param temperature What sampling temperature to use, between 0 and 1. * @param temperature What sampling temperature to use, between 0 and 1.
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Float temperature) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Double temperature) {
this(messages, model, null, null, null, null, null, null, null, this(messages, model, null, null, null, null, null, null, null,
null, null, null, false, null, temperature, null, null, null, null, false, null, temperature, null,
null, null, null, null); null, null, null, null);
@@ -456,7 +456,7 @@ public class OpenAiApi {
* @param stream If set, partial message deltas will be sent.Tokens will be sent as data-only server-sent events * @param stream If set, partial message deltas will be sent.Tokens will be sent as data-only server-sent events
* as they become available, with the stream terminated by a data: [DONE] message. * as they become available, with the stream terminated by a data: [DONE] message.
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Float temperature, boolean stream) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Double temperature, boolean stream) {
this(messages, model, null, null, null, null, null, null, null, this(messages, model, null, null, null, null, null, null, null,
null, null, null, stream, null, temperature, null, null, null, null, stream, null, temperature, null,
null, null, null, null); null, null, null, null);
@@ -474,7 +474,7 @@ public class OpenAiApi {
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model,
List<FunctionTool> tools, Object toolChoice) { List<FunctionTool> tools, Object toolChoice) {
this(messages, model, null, null, null, null, null, null, null, this(messages, model, null, null, null, null, null, null, null,
null, null, null, false, null, 0.8f, null, null, null, null, false, null, 0.8, null,
tools, toolChoice, null, null); tools, toolChoice, null, null);
} }

View File

@@ -35,7 +35,7 @@ public class ChatCompletionRequestTests {
public void createRequestWithChatOptions() { public void createRequestWithChatOptions() {
var client = new OpenAiChatModel(new OpenAiApi("TEST"), var client = new OpenAiChatModel(new OpenAiApi("TEST"),
OpenAiChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6f).build()); OpenAiChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6).build());
var request = client.createRequest(new Prompt("Test message content"), false); var request = client.createRequest(new Prompt("Test message content"), false);
@@ -43,16 +43,16 @@ public class ChatCompletionRequestTests {
assertThat(request.stream()).isFalse(); assertThat(request.stream()).isFalse();
assertThat(request.model()).isEqualTo("DEFAULT_MODEL"); assertThat(request.model()).isEqualTo("DEFAULT_MODEL");
assertThat(request.temperature()).isEqualTo(66.6f); assertThat(request.temperature()).isEqualTo(66.6);
request = client.createRequest(new Prompt("Test message content", request = client.createRequest(new Prompt("Test message content",
OpenAiChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9f).build()), true); OpenAiChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9).build()), true);
assertThat(request.messages()).hasSize(1); assertThat(request.messages()).hasSize(1);
assertThat(request.stream()).isTrue(); assertThat(request.stream()).isTrue();
assertThat(request.model()).isEqualTo("PROMPT_MODEL"); assertThat(request.model()).isEqualTo("PROMPT_MODEL");
assertThat(request.temperature()).isEqualTo(99.9f); assertThat(request.temperature()).isEqualTo(99.9);
} }
@Test @Test

View File

@@ -44,7 +44,7 @@ public class OpenAiApiIT {
void chatCompletionEntity() { void chatCompletionEntity() {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER); ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER);
ResponseEntity<ChatCompletion> response = openAiApi.chatCompletionEntity( ResponseEntity<ChatCompletion> response = openAiApi.chatCompletionEntity(
new ChatCompletionRequest(List.of(chatCompletionMessage), "gpt-3.5-turbo", 0.8f, false)); new ChatCompletionRequest(List.of(chatCompletionMessage), "gpt-3.5-turbo", 0.8, false));
assertThat(response).isNotNull(); assertThat(response).isNotNull();
assertThat(response.getBody()).isNotNull(); assertThat(response.getBody()).isNotNull();
@@ -54,7 +54,7 @@ public class OpenAiApiIT {
void chatCompletionStream() { void chatCompletionStream() {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER); ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER);
Flux<ChatCompletionChunk> response = openAiApi.chatCompletionStream( Flux<ChatCompletionChunk> response = openAiApi.chatCompletionStream(
new ChatCompletionRequest(List.of(chatCompletionMessage), "gpt-3.5-turbo", 0.8f, true)); new ChatCompletionRequest(List.of(chatCompletionMessage), "gpt-3.5-turbo", 0.8, true));
assertThat(response).isNotNull(); assertThat(response).isNotNull();
assertThat(response.collectList().block()).isNotNull(); assertThat(response.collectList().block()).isNotNull();

View File

@@ -124,7 +124,7 @@ public class OpenAiApiToolFunctionCallIT {
} }
} }
var functionResponseRequest = new ChatCompletionRequest(messages, "gpt-4o", 0.5f); var functionResponseRequest = new ChatCompletionRequest(messages, "gpt-4o", 0.5);
ResponseEntity<ChatCompletion> chatCompletion2 = completionApi.chatCompletionEntity(functionResponseRequest); ResponseEntity<ChatCompletion> chatCompletion2 = completionApi.chatCompletionEntity(functionResponseRequest);

View File

@@ -71,12 +71,12 @@ public class OpenAiChatModelObservationIT {
var options = OpenAiChatOptions.builder() var options = OpenAiChatOptions.builder()
.withModel(OpenAiApi.ChatModel.GPT_4_O_MINI.getValue()) .withModel(OpenAiApi.ChatModel.GPT_4_O_MINI.getValue())
.withFrequencyPenalty(0f) .withFrequencyPenalty(0.0)
.withMaxTokens(2048) .withMaxTokens(2048)
.withPresencePenalty(0f) .withPresencePenalty(0.0)
.withStop(List.of("this-is-the-end")) .withStop(List.of("this-is-the-end"))
.withTemperature(0.7f) .withTemperature(0.7)
.withTopP(1f) .withTopP(1.0)
.build(); .build();
Prompt prompt = new Prompt("Why does a raven look like a desk?", options); Prompt prompt = new Prompt("Why does a raven look like a desk?", options);
@@ -94,12 +94,12 @@ public class OpenAiChatModelObservationIT {
void observationForStreamingChatOperation() { void observationForStreamingChatOperation() {
var options = OpenAiChatOptions.builder() var options = OpenAiChatOptions.builder()
.withModel(OpenAiApi.ChatModel.GPT_4_O_MINI.getValue()) .withModel(OpenAiApi.ChatModel.GPT_4_O_MINI.getValue())
.withFrequencyPenalty(0f) .withFrequencyPenalty(0.0)
.withMaxTokens(2048) .withMaxTokens(2048)
.withPresencePenalty(0f) .withPresencePenalty(0.0)
.withStop(List.of("this-is-the-end")) .withStop(List.of("this-is-the-end"))
.withTemperature(0.7f) .withTemperature(0.7)
.withTopP(1f) .withTopP(1.0)
.withStreamUsage(true) .withStreamUsage(true)
.build(); .build();

View File

@@ -203,7 +203,7 @@ public class OpenAiPaymentTransactionIT {
return new OpenAiChatModel(openAiApi, return new OpenAiChatModel(openAiApi,
OpenAiChatOptions.builder() OpenAiChatOptions.builder()
.withModel(ChatModel.GPT_4_O_MINI.getName()) .withModel(ChatModel.GPT_4_O_MINI.getName())
.withTemperature(0.1f) .withTemperature(0.1)
.build(), .build(),
functionCallbackContext, RetryUtils.DEFAULT_RETRY_TEMPLATE); functionCallbackContext, RetryUtils.DEFAULT_RETRY_TEMPLATE);
} }

View File

@@ -77,7 +77,7 @@ public class QianFanChatModel implements ChatModel, StreamingChatModel {
*/ */
public QianFanChatModel(QianFanApi qianFanApi) { public QianFanChatModel(QianFanApi qianFanApi) {
this(qianFanApi, this(qianFanApi,
QianFanChatOptions.builder().withModel(QianFanApi.DEFAULT_CHAT_MODEL).withTemperature(0.7f).build()); QianFanChatOptions.builder().withModel(QianFanApi.DEFAULT_CHAT_MODEL).withTemperature(0.7).build());
} }
/** /**

View File

@@ -46,7 +46,7 @@ public class QianFanChatOptions implements ChatOptions {
* Number between -2.0 and 2.0. Positive values penalize new tokens based on their existing * Number between -2.0 and 2.0. Positive values penalize new tokens based on their existing
* frequency in the text so far, decreasing the model's likelihood to repeat the same line verbatim. * frequency in the text so far, decreasing the model's likelihood to repeat the same line verbatim.
*/ */
private @JsonProperty("frequency_penalty") Float frequencyPenalty; private @JsonProperty("frequency_penalty") Double frequencyPenalty;
/** /**
* The maximum number of tokens to generate in the chat completion. The total length of input * The maximum number of tokens to generate in the chat completion. The total length of input
* tokens and generated tokens is limited by the model's context length. * tokens and generated tokens is limited by the model's context length.
@@ -56,7 +56,7 @@ public class QianFanChatOptions implements ChatOptions {
* Number between -2.0 and 2.0. Positive values penalize new tokens based on whether they * Number between -2.0 and 2.0. Positive values penalize new tokens based on whether they
* appear in the text so far, increasing the model's likelihood to talk about new topics. * appear in the text so far, increasing the model's likelihood to talk about new topics.
*/ */
private @JsonProperty("presence_penalty") Float presencePenalty; private @JsonProperty("presence_penalty") Double presencePenalty;
/** /**
* An object specifying the format that the model must output. Setting to { "type": * An object specifying the format that the model must output. Setting to { "type":
* "json_object" } enables JSON mode, which guarantees the message the model generates is valid JSON. * "json_object" } enables JSON mode, which guarantees the message the model generates is valid JSON.
@@ -72,13 +72,13 @@ public class QianFanChatOptions implements ChatOptions {
* more random, while lower values like 0.2 will make it more focused and deterministic. We generally recommend * more random, while lower values like 0.2 will make it more focused and deterministic. We generally recommend
* altering this or top_p but not both. * altering this or top_p but not both.
*/ */
private @JsonProperty("temperature") Float temperature; private @JsonProperty("temperature") Double temperature;
/** /**
* An alternative to sampling with temperature, called nucleus sampling, where the model considers the * An alternative to sampling with temperature, called nucleus sampling, where the model considers the
* results of the tokens with top_p probability mass. So 0.1 means only the tokens comprising the top 10% * results of the tokens with top_p probability mass. So 0.1 means only the tokens comprising the top 10%
* probability mass are considered. We generally recommend altering this or temperature but not both. * probability mass are considered. We generally recommend altering this or temperature but not both.
*/ */
private @JsonProperty("top_p") Float topP; private @JsonProperty("top_p") Double topP;
// @formatter:on // @formatter:on
public static Builder builder() { public static Builder builder() {
@@ -102,7 +102,7 @@ public class QianFanChatOptions implements ChatOptions {
return this; return this;
} }
public Builder withFrequencyPenalty(Float frequencyPenalty) { public Builder withFrequencyPenalty(Double frequencyPenalty) {
this.options.frequencyPenalty = frequencyPenalty; this.options.frequencyPenalty = frequencyPenalty;
return this; return this;
} }
@@ -112,7 +112,7 @@ public class QianFanChatOptions implements ChatOptions {
return this; return this;
} }
public Builder withPresencePenalty(Float presencePenalty) { public Builder withPresencePenalty(Double presencePenalty) {
this.options.presencePenalty = presencePenalty; this.options.presencePenalty = presencePenalty;
return this; return this;
} }
@@ -127,12 +127,12 @@ public class QianFanChatOptions implements ChatOptions {
return this; return this;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.temperature = temperature; this.options.temperature = temperature;
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.topP = topP; this.options.topP = topP;
return this; return this;
} }
@@ -153,11 +153,11 @@ public class QianFanChatOptions implements ChatOptions {
} }
@Override @Override
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return this.frequencyPenalty; return this.frequencyPenalty;
} }
public void setFrequencyPenalty(Float frequencyPenalty) { public void setFrequencyPenalty(Double frequencyPenalty) {
this.frequencyPenalty = frequencyPenalty; this.frequencyPenalty = frequencyPenalty;
} }
@@ -171,11 +171,11 @@ public class QianFanChatOptions implements ChatOptions {
} }
@Override @Override
public Float getPresencePenalty() { public Double getPresencePenalty() {
return this.presencePenalty; return this.presencePenalty;
} }
public void setPresencePenalty(Float presencePenalty) { public void setPresencePenalty(Double presencePenalty) {
this.presencePenalty = presencePenalty; this.presencePenalty = presencePenalty;
} }
@@ -207,20 +207,20 @@ public class QianFanChatOptions implements ChatOptions {
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }

View File

@@ -39,6 +39,7 @@ import java.util.function.Predicate;
* <a href="https://cloud.baidu.com/doc/WENXINWORKSHOP/index.html">QianFan Docs</a> * <a href="https://cloud.baidu.com/doc/WENXINWORKSHOP/index.html">QianFan Docs</a>
* *
* @author Geng Rong * @author Geng Rong
* @author Thomas Vitale
* @since 1.0 * @since 1.0
*/ */
public class QianFanApi extends AuthApi { public class QianFanApi extends AuthApi {
@@ -187,14 +188,14 @@ public class QianFanApi extends AuthApi {
@JsonProperty("messages") List<ChatCompletionMessage> messages, @JsonProperty("messages") List<ChatCompletionMessage> messages,
@JsonProperty("system") String system, @JsonProperty("system") String system,
@JsonProperty("model") String model, @JsonProperty("model") String model,
@JsonProperty("frequency_penalty") Float frequencyPenalty, @JsonProperty("frequency_penalty") Double frequencyPenalty,
@JsonProperty("max_output_tokens") Integer maxTokens, @JsonProperty("max_output_tokens") Integer maxTokens,
@JsonProperty("presence_penalty") Float presencePenalty, @JsonProperty("presence_penalty") Double presencePenalty,
@JsonProperty("response_format") ResponseFormat responseFormat, @JsonProperty("response_format") ResponseFormat responseFormat,
@JsonProperty("stop") List<String> stop, @JsonProperty("stop") List<String> stop,
@JsonProperty("stream") Boolean stream, @JsonProperty("stream") Boolean stream,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("top_p") Float topP) { @JsonProperty("top_p") Double topP) {
/** /**
* Shortcut constructor for a chat completion request with the given messages and model. * Shortcut constructor for a chat completion request with the given messages and model.
@@ -203,7 +204,7 @@ public class QianFanApi extends AuthApi {
* @param model ID of the model to use. * @param model ID of the model to use.
* @param temperature What sampling temperature to use, between 0 and 1. * @param temperature What sampling temperature to use, between 0 and 1.
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String system, String model, Float temperature) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, String system, String model, Double temperature) {
this(messages, system, model, null,null, this(messages, system, model, null,null,
null, null, null, false, temperature, null); null, null, null, false, temperature, null);
} }
@@ -217,7 +218,7 @@ public class QianFanApi extends AuthApi {
* @param stream If set, partial message deltas will be sent.Tokens will be sent as data-only server-sent events * @param stream If set, partial message deltas will be sent.Tokens will be sent as data-only server-sent events
* as they become available, with the stream terminated by a data: [DONE] message. * as they become available, with the stream terminated by a data: [DONE] message.
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String system, String model, Float temperature, boolean stream) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, String system, String model, Double temperature, boolean stream) {
this(messages, system, model, null,null, this(messages, system, model, null,null,
null, null, null, stream, temperature, null); null, null, null, stream, temperature, null);
} }
@@ -233,7 +234,7 @@ public class QianFanApi extends AuthApi {
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String system, Boolean stream) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, String system, Boolean stream) {
this(messages, system, DEFAULT_CHAT_MODEL, null,null, this(messages, system, DEFAULT_CHAT_MODEL, null,null,
null, null, null, stream, 0.8F, null); null, null, null, stream, 0.8, null);
} }
/** /**

View File

@@ -30,7 +30,7 @@ public class ChatCompletionRequestTests {
public void createRequestWithChatOptions() { public void createRequestWithChatOptions() {
var client = new QianFanChatModel(new QianFanApi("TEST", "TEST"), var client = new QianFanChatModel(new QianFanApi("TEST", "TEST"),
QianFanChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6f).build()); QianFanChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6).build());
var request = client.createRequest(new Prompt("Test message content"), false); var request = client.createRequest(new Prompt("Test message content"), false);
@@ -38,16 +38,16 @@ public class ChatCompletionRequestTests {
assertThat(request.stream()).isFalse(); assertThat(request.stream()).isFalse();
assertThat(request.model()).isEqualTo("DEFAULT_MODEL"); assertThat(request.model()).isEqualTo("DEFAULT_MODEL");
assertThat(request.temperature()).isEqualTo(66.6f); assertThat(request.temperature()).isEqualTo(66.6);
request = client.createRequest(new Prompt("Test message content", request = client.createRequest(new Prompt("Test message content",
QianFanChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9f).build()), true); QianFanChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9).build()), true);
assertThat(request.messages()).hasSize(1); assertThat(request.messages()).hasSize(1);
assertThat(request.stream()).isTrue(); assertThat(request.stream()).isTrue();
assertThat(request.model()).isEqualTo("PROMPT_MODEL"); assertThat(request.model()).isEqualTo("PROMPT_MODEL");
assertThat(request.temperature()).isEqualTo(99.9f); assertThat(request.temperature()).isEqualTo(99.9);
} }
} }

View File

@@ -47,7 +47,7 @@ public class QianFanApiIT {
void chatCompletionEntity() { void chatCompletionEntity() {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER); ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER);
ResponseEntity<ChatCompletion> response = qianFanApi.chatCompletionEntity(new ChatCompletionRequest( ResponseEntity<ChatCompletion> response = qianFanApi.chatCompletionEntity(new ChatCompletionRequest(
List.of(chatCompletionMessage), buildSystemMessage(), "ernie_speed", 0.7f, false)); List.of(chatCompletionMessage), buildSystemMessage(), "ernie_speed", 0.7, false));
assertThat(response).isNotNull(); assertThat(response).isNotNull();
assertThat(response.getBody()).isNotNull(); assertThat(response.getBody()).isNotNull();
@@ -57,7 +57,7 @@ public class QianFanApiIT {
void chatCompletionStream() { void chatCompletionStream() {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER); ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER);
Flux<ChatCompletionChunk> response = qianFanApi.chatCompletionStream(new ChatCompletionRequest( Flux<ChatCompletionChunk> response = qianFanApi.chatCompletionStream(new ChatCompletionRequest(
List.of(chatCompletionMessage), buildSystemMessage(), "ernie_speed", 0.7f, true)); List.of(chatCompletionMessage), buildSystemMessage(), "ernie_speed", 0.7, true));
assertThat(response).isNotNull(); assertThat(response).isNotNull();
assertThat(response.collectList().block()).isNotNull(); assertThat(response.collectList().block()).isNotNull();

View File

@@ -138,7 +138,7 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
public VertexAiGeminiChatModel(VertexAI vertexAI) { public VertexAiGeminiChatModel(VertexAI vertexAI) {
this(vertexAI, this(vertexAI,
VertexAiGeminiChatOptions.builder().withModel(ChatModel.GEMINI_1_5_PRO).withTemperature(0.8f).build()); VertexAiGeminiChatOptions.builder().withModel(ChatModel.GEMINI_1_5_PRO).withTemperature(0.8).build());
} }
public VertexAiGeminiChatModel(VertexAI vertexAI, VertexAiGeminiChatOptions options) { public VertexAiGeminiChatModel(VertexAI vertexAI, VertexAiGeminiChatOptions options) {
@@ -354,7 +354,7 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
GenerationConfig.Builder generationConfigBuilder = GenerationConfig.newBuilder(); GenerationConfig.Builder generationConfigBuilder = GenerationConfig.newBuilder();
if (options.getTemperature() != null) { if (options.getTemperature() != null) {
generationConfigBuilder.setTemperature(options.getTemperature()); generationConfigBuilder.setTemperature(options.getTemperature().floatValue());
} }
if (options.getMaxOutputTokens() != null) { if (options.getMaxOutputTokens() != null) {
generationConfigBuilder.setMaxOutputTokens(options.getMaxOutputTokens()); generationConfigBuilder.setMaxOutputTokens(options.getMaxOutputTokens());
@@ -363,7 +363,7 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
generationConfigBuilder.setTopK(options.getTopK()); generationConfigBuilder.setTopK(options.getTopK());
} }
if (options.getTopP() != null) { if (options.getTopP() != null) {
generationConfigBuilder.setTopP(options.getTopP()); generationConfigBuilder.setTopP(options.getTopP().floatValue());
} }
if (options.getCandidateCount() != null) { if (options.getCandidateCount() != null) {
generationConfigBuilder.setCandidateCount(options.getCandidateCount()); generationConfigBuilder.setCandidateCount(options.getCandidateCount());

View File

@@ -58,11 +58,11 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions, ChatOp
/** /**
* Optional. Controls the randomness of predictions. * Optional. Controls the randomness of predictions.
*/ */
private @JsonProperty("temperature") Float temperature; private @JsonProperty("temperature") Double temperature;
/** /**
* Optional. If specified, nucleus sampling will be used. * Optional. If specified, nucleus sampling will be used.
*/ */
private @JsonProperty("topP") Float topP; private @JsonProperty("topP") Double topP;
/** /**
* Optional. If specified, top k sampling will be used. * Optional. If specified, top k sampling will be used.
*/ */
@@ -131,12 +131,12 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions, ChatOp
return this; return this;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.setTemperature(temperature); this.options.setTemperature(temperature);
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.setTopP(topP); this.options.setTopP(topP);
return this; return this;
} }
@@ -210,20 +210,20 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions, ChatOp
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }
@@ -303,13 +303,13 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions, ChatOp
@Override @Override
@JsonIgnore @JsonIgnore
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return null; return null;
} }
@Override @Override
@JsonIgnore @JsonIgnore
public Float getPresencePenalty() { public Double getPresencePenalty() {
return null; return null;
} }

View File

@@ -51,7 +51,7 @@ public class CreateGeminiRequestTests {
public void createRequestWithChatOptions() { public void createRequestWithChatOptions() {
var client = new VertexAiGeminiChatModel(vertexAI, var client = new VertexAiGeminiChatModel(vertexAI,
VertexAiGeminiChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6f).build()); VertexAiGeminiChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6).build());
GeminiRequest request = client.createGeminiRequest(new Prompt("Test message content")); GeminiRequest request = client.createGeminiRequest(new Prompt("Test message content"));
@@ -62,7 +62,7 @@ public class CreateGeminiRequestTests {
assertThat(request.model().getGenerationConfig().getTemperature()).isEqualTo(66.6f); assertThat(request.model().getGenerationConfig().getTemperature()).isEqualTo(66.6f);
request = client.createGeminiRequest(new Prompt("Test message content", request = client.createGeminiRequest(new Prompt("Test message content",
VertexAiGeminiChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9f).build())); VertexAiGeminiChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9).build()));
assertThat(request.contents()).hasSize(1); assertThat(request.contents()).hasSize(1);
@@ -80,7 +80,7 @@ public class CreateGeminiRequestTests {
List.of(new Media(MimeTypeUtils.IMAGE_PNG, new URL("http://example.com")))); List.of(new Media(MimeTypeUtils.IMAGE_PNG, new URL("http://example.com"))));
var client = new VertexAiGeminiChatModel(vertexAI, var client = new VertexAiGeminiChatModel(vertexAI,
VertexAiGeminiChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6f).build()); VertexAiGeminiChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6).build());
GeminiRequest request = client.createGeminiRequest(new Prompt(List.of(systemMessage, userMessage))); GeminiRequest request = client.createGeminiRequest(new Prompt(List.of(systemMessage, userMessage)));
@@ -197,10 +197,10 @@ public class CreateGeminiRequestTests {
var client = new VertexAiGeminiChatModel(vertexAI, var client = new VertexAiGeminiChatModel(vertexAI,
VertexAiGeminiChatOptions.builder() VertexAiGeminiChatOptions.builder()
.withModel("DEFAULT_MODEL") .withModel("DEFAULT_MODEL")
.withTemperature(66.6f) .withTemperature(66.6)
.withMaxOutputTokens(100) .withMaxOutputTokens(100)
.withTopK(10.0f) .withTopK(10.0f)
.withTopP(5.0f) .withTopP(5.0)
.withStopSequences(List.of("stop1", "stop2")) .withStopSequences(List.of("stop1", "stop2"))
.withCandidateCount(1) .withCandidateCount(1)
.withResponseMimeType("application/json") .withResponseMimeType("application/json")

View File

@@ -242,7 +242,7 @@ public class VertexAiGeminiChatModelFunctionCallingIT {
return new VertexAiGeminiChatModel(vertexAi, return new VertexAiGeminiChatModel(vertexAi,
VertexAiGeminiChatOptions.builder() VertexAiGeminiChatOptions.builder()
.withModel(VertexAiGeminiChatModel.ChatModel.GEMINI_1_5_PRO) .withModel(VertexAiGeminiChatModel.ChatModel.GEMINI_1_5_PRO)
.withTemperature(0.9f) .withTemperature(0.9)
.build()); .build());
} }

View File

@@ -202,7 +202,7 @@ public class VertexAiGeminiPaymentTransactionIT {
return new VertexAiGeminiChatModel(vertexAi, return new VertexAiGeminiChatModel(vertexAi,
VertexAiGeminiChatOptions.builder() VertexAiGeminiChatOptions.builder()
.withModel(VertexAiGeminiChatModel.ChatModel.GEMINI_1_5_FLASH) .withModel(VertexAiGeminiChatModel.ChatModel.GEMINI_1_5_FLASH)
.withTemperature(0.1f) .withTemperature(0.1)
.build(), .build(),
functionCallbackContext); functionCallbackContext);
} }

View File

@@ -43,7 +43,7 @@ public class VertexAiPaLm2ChatModel implements ChatModel {
public VertexAiPaLm2ChatModel(VertexAiPaLm2Api vertexAiApi) { public VertexAiPaLm2ChatModel(VertexAiPaLm2Api vertexAiApi) {
this(vertexAiApi, this(vertexAiApi,
VertexAiPaLm2ChatOptions.builder().withTemperature(0.7f).withCandidateCount(1).withTopK(20).build()); VertexAiPaLm2ChatOptions.builder().withTemperature(0.7).withCandidateCount(1).withTopK(20).build());
} }
public VertexAiPaLm2ChatModel(VertexAiPaLm2Api vertexAiApi, VertexAiPaLm2ChatOptions defaultOptions) { public VertexAiPaLm2ChatModel(VertexAiPaLm2Api vertexAiApi, VertexAiPaLm2ChatOptions defaultOptions) {

View File

@@ -39,7 +39,7 @@ public class VertexAiPaLm2ChatOptions implements ChatOptions {
* generative. This value specifies default to be used by the backend while making the * generative. This value specifies default to be used by the backend while making the
* call to the generative. * call to the generative.
*/ */
private @JsonProperty("temperature") Float temperature; private @JsonProperty("temperature") Double temperature;
/** /**
* The number of generated response messages to return. This value must be between [1, * The number of generated response messages to return. This value must be between [1,
@@ -52,7 +52,7 @@ public class VertexAiPaLm2ChatOptions implements ChatOptions {
* generative uses combined Top-k and nucleus sampling. Nucleus sampling considers the * generative uses combined Top-k and nucleus sampling. Nucleus sampling considers the
* smallest set of tokens whose probability sum is at least topP. * smallest set of tokens whose probability sum is at least topP.
*/ */
private @JsonProperty("topP") Float topP; private @JsonProperty("topP") Double topP;
/** /**
* The maximum number of tokens to consider when sampling. The generative uses * The maximum number of tokens to consider when sampling. The generative uses
@@ -70,7 +70,7 @@ public class VertexAiPaLm2ChatOptions implements ChatOptions {
private VertexAiPaLm2ChatOptions options = new VertexAiPaLm2ChatOptions(); private VertexAiPaLm2ChatOptions options = new VertexAiPaLm2ChatOptions();
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.temperature = temperature; this.options.temperature = temperature;
return this; return this;
} }
@@ -80,7 +80,7 @@ public class VertexAiPaLm2ChatOptions implements ChatOptions {
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.topP = topP; this.options.topP = topP;
return this; return this;
} }
@@ -97,11 +97,11 @@ public class VertexAiPaLm2ChatOptions implements ChatOptions {
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@@ -114,11 +114,11 @@ public class VertexAiPaLm2ChatOptions implements ChatOptions {
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }
@@ -151,13 +151,13 @@ public class VertexAiPaLm2ChatOptions implements ChatOptions {
@Override @Override
@JsonIgnore @JsonIgnore
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return null; return null;
} }
@Override @Override
@JsonIgnore @JsonIgnore
public Float getPresencePenalty() { public Double getPresencePenalty() {
return null; return null;
} }

View File

@@ -87,6 +87,7 @@ import org.springframework.web.client.RestClient;
* https://ai.google.dev/api/rest#rest-resource:-v1.models * https://ai.google.dev/api/rest#rest-resource:-v1.models
* *
* @author Christian Tzolov * @author Christian Tzolov
* @author Thomas Vitale
*/ */
public class VertexAiPaLm2Api { public class VertexAiPaLm2Api {
@@ -353,8 +354,8 @@ public class VertexAiPaLm2Api {
@JsonProperty("inputTokenLimit") Integer inputTokenLimit, @JsonProperty("inputTokenLimit") Integer inputTokenLimit,
@JsonProperty("outputTokenLimit") Integer outputTokenLimit, @JsonProperty("outputTokenLimit") Integer outputTokenLimit,
@JsonProperty("supportedGenerationMethods") List<String> supportedGenerationMethods, @JsonProperty("supportedGenerationMethods") List<String> supportedGenerationMethods,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("topP") Float topP, @JsonProperty("topP") Double topP,
@JsonProperty("topK") Integer topK) { @JsonProperty("topK") Integer topK) {
} }
@@ -528,9 +529,9 @@ public class VertexAiPaLm2Api {
@JsonInclude(Include.NON_NULL) @JsonInclude(Include.NON_NULL)
public record GenerateMessageRequest( public record GenerateMessageRequest(
@JsonProperty("prompt") MessagePrompt prompt, @JsonProperty("prompt") MessagePrompt prompt,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("candidateCount") Integer candidateCount, @JsonProperty("candidateCount") Integer candidateCount,
@JsonProperty("topP") Float topP, @JsonProperty("topP") Double topP,
@JsonProperty("topK") Integer topK) { @JsonProperty("topK") Integer topK) {
/** /**
@@ -549,7 +550,7 @@ public class VertexAiPaLm2Api {
* @param temperature (optional) Controls the randomness of the output. * @param temperature (optional) Controls the randomness of the output.
* @param topK (optional) The maximum number of tokens to consider when sampling. * @param topK (optional) The maximum number of tokens to consider when sampling.
*/ */
public GenerateMessageRequest(MessagePrompt prompt, Float temperature, Integer topK) { public GenerateMessageRequest(MessagePrompt prompt, Double temperature, Integer topK) {
this(prompt, temperature, null, null, topK); this(prompt, temperature, null, null, topK);
} }
} }

View File

@@ -39,7 +39,7 @@ public class VertexAiPaLm2ChatRequestTests {
assertThat(request.prompt().messages()).hasSize(1); assertThat(request.prompt().messages()).hasSize(1);
assertThat(request.candidateCount()).isEqualTo(1); assertThat(request.candidateCount()).isEqualTo(1);
assertThat(request.temperature()).isEqualTo(0.7f); assertThat(request.temperature()).isEqualTo(0.7);
assertThat(request.topK()).isEqualTo(20); assertThat(request.topK()).isEqualTo(20);
assertThat(request.topP()).isNull(); assertThat(request.topP()).isNull();
} }
@@ -49,8 +49,8 @@ public class VertexAiPaLm2ChatRequestTests {
// Runtime options should override the default options. // Runtime options should override the default options.
VertexAiPaLm2ChatOptions promptOptions = VertexAiPaLm2ChatOptions.builder() VertexAiPaLm2ChatOptions promptOptions = VertexAiPaLm2ChatOptions.builder()
.withTemperature(0.8f) .withTemperature(0.8)
.withTopP(0.5f) .withTopP(0.5)
.withTopK(99) .withTopK(99)
// .withCandidateCount(2) // .withCandidateCount(2)
.build(); .build();
@@ -60,9 +60,9 @@ public class VertexAiPaLm2ChatRequestTests {
assertThat(request.prompt().messages()).hasSize(1); assertThat(request.prompt().messages()).hasSize(1);
assertThat(request.candidateCount()).isEqualTo(1); assertThat(request.candidateCount()).isEqualTo(1);
assertThat(request.temperature()).isEqualTo(0.8f); assertThat(request.temperature()).isEqualTo(0.8);
assertThat(request.topK()).isEqualTo(99); assertThat(request.topK()).isEqualTo(99);
assertThat(request.topP()).isEqualTo(0.5f); assertThat(request.topP()).isEqualTo(0.5);
} }
@Test @Test
@@ -70,9 +70,9 @@ public class VertexAiPaLm2ChatRequestTests {
// runtime options. // runtime options.
ChatOptions portablePromptOptions = ChatOptionsBuilder.builder() ChatOptions portablePromptOptions = ChatOptionsBuilder.builder()
.withTemperature(0.9f) .withTemperature(0.9)
.withTopK(100) .withTopK(100)
.withTopP(0.6f) .withTopP(0.6)
.build(); .build();
var request = chatModel.createRequest(new Prompt("Test message content", portablePromptOptions)); var request = chatModel.createRequest(new Prompt("Test message content", portablePromptOptions));
@@ -80,9 +80,9 @@ public class VertexAiPaLm2ChatRequestTests {
assertThat(request.prompt().messages()).hasSize(1); assertThat(request.prompt().messages()).hasSize(1);
assertThat(request.candidateCount()).isEqualTo(1); assertThat(request.candidateCount()).isEqualTo(1);
assertThat(request.temperature()).isEqualTo(0.9f); assertThat(request.temperature()).isEqualTo(0.9);
assertThat(request.topK()).isEqualTo(100); assertThat(request.topK()).isEqualTo(100);
assertThat(request.topP()).isEqualTo(0.6f); assertThat(request.topP()).isEqualTo(0.6);
} }
} }

View File

@@ -57,13 +57,13 @@ public class WatsonxAiChatModel implements ChatModel, StreamingChatModel {
public WatsonxAiChatModel(WatsonxAiApi watsonxAiApi) { public WatsonxAiChatModel(WatsonxAiApi watsonxAiApi) {
this(watsonxAiApi, this(watsonxAiApi,
WatsonxAiChatOptions.builder() WatsonxAiChatOptions.builder()
.withTemperature(0.7f) .withTemperature(0.7)
.withTopP(1.0f) .withTopP(1.0)
.withTopK(50) .withTopK(50)
.withDecodingMethod("greedy") .withDecodingMethod("greedy")
.withMaxNewTokens(20) .withMaxNewTokens(20)
.withMinNewTokens(0) .withMinNewTokens(0)
.withRepetitionPenalty(1.0f) .withRepetitionPenalty(1.0)
.withStopSequences(List.of()) .withStopSequences(List.of())
.build()); .build());
} }

View File

@@ -48,14 +48,14 @@ public class WatsonxAiChatOptions implements ChatOptions {
* The temperature of the model. Increasing the temperature will * The temperature of the model. Increasing the temperature will
* make the model answer more creatively. (Default: 0.7) * make the model answer more creatively. (Default: 0.7)
*/ */
@JsonProperty("temperature") private Float temperature; @JsonProperty("temperature") private Double temperature;
/** /**
* Works together with top-k. A higher value (e.g., 0.95) will lead to * Works together with top-k. A higher value (e.g., 0.95) will lead to
* more diverse text, while a lower value (e.g., 0.2) will generate more focused and * more diverse text, while a lower value (e.g., 0.2) will generate more focused and
* conservative text. (Default: 1.0) * conservative text. (Default: 1.0)
*/ */
@JsonProperty("top_p") private Float topP; @JsonProperty("top_p") private Double topP;
/** /**
* Reduces the probability of generating nonsense. A higher value (e.g. * Reduces the probability of generating nonsense. A higher value (e.g.
@@ -104,7 +104,7 @@ public class WatsonxAiChatOptions implements ChatOptions {
* (e.g., 1.8) will penalize repetitions more strongly, while a lower value (e.g., * (e.g., 1.8) will penalize repetitions more strongly, while a lower value (e.g.,
* 1.1) will be more lenient. (Default: 1.0) * 1.1) will be more lenient. (Default: 1.0)
*/ */
@JsonProperty("repetition_penalty") private Float repetitionPenalty; @JsonProperty("repetition_penalty") private Double repetitionPenalty;
/** /**
* Produce repeatable results, set the same random seed value every time. (Default: randomly generated) * Produce repeatable results, set the same random seed value every time. (Default: randomly generated)
@@ -126,20 +126,20 @@ public class WatsonxAiChatOptions implements ChatOptions {
private ObjectMapper mapper = new ObjectMapper(); private ObjectMapper mapper = new ObjectMapper();
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return temperature; return temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return topP; return topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }
@@ -198,20 +198,20 @@ public class WatsonxAiChatOptions implements ChatOptions {
@Override @Override
@JsonIgnore @JsonIgnore
public Float getPresencePenalty() { public Double getPresencePenalty() {
return getRepetitionPenalty(); return getRepetitionPenalty();
} }
@JsonIgnore @JsonIgnore
public void setPresencePenalty(Float presencePenalty) { public void setPresencePenalty(Double presencePenalty) {
setRepetitionPenalty(presencePenalty); setRepetitionPenalty(presencePenalty);
} }
public Float getRepetitionPenalty() { public Double getRepetitionPenalty() {
return repetitionPenalty; return repetitionPenalty;
} }
public void setRepetitionPenalty(Float repetitionPenalty) { public void setRepetitionPenalty(Double repetitionPenalty) {
this.repetitionPenalty = repetitionPenalty; this.repetitionPenalty = repetitionPenalty;
} }
@@ -248,7 +248,7 @@ public class WatsonxAiChatOptions implements ChatOptions {
@Override @Override
@JsonIgnore @JsonIgnore
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return null; return null;
} }
@@ -260,12 +260,12 @@ public class WatsonxAiChatOptions implements ChatOptions {
WatsonxAiChatOptions options = new WatsonxAiChatOptions(); WatsonxAiChatOptions options = new WatsonxAiChatOptions();
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.temperature = temperature; this.options.temperature = temperature;
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.topP = topP; this.options.topP = topP;
return this; return this;
} }
@@ -295,7 +295,7 @@ public class WatsonxAiChatOptions implements ChatOptions {
return this; return this;
} }
public Builder withRepetitionPenalty(Float repetitionPenalty) { public Builder withRepetitionPenalty(Double repetitionPenalty) {
this.options.repetitionPenalty = repetitionPenalty; this.options.repetitionPenalty = repetitionPenalty;
return this; return this;
} }

View File

@@ -52,7 +52,7 @@ public class WatsonxAiChatModelTest {
@Test @Test
public void testCreateRequestWithNoModelId() { public void testCreateRequestWithNoModelId() {
var options = ChatOptionsBuilder.builder().withTemperature(0.9f).withTopK(100).withTopP(0.6f).build(); var options = ChatOptionsBuilder.builder().withTemperature(0.9).withTopK(100).withTopP(0.6).build();
Prompt prompt = new Prompt("Test message", options); Prompt prompt = new Prompt("Test message", options);
@@ -93,12 +93,12 @@ public class WatsonxAiChatModelTest {
WatsonxAiChatOptions modelOptions = WatsonxAiChatOptions.builder() WatsonxAiChatOptions modelOptions = WatsonxAiChatOptions.builder()
.withModel("meta-llama/llama-2-70b-chat") .withModel("meta-llama/llama-2-70b-chat")
.withDecodingMethod("sample") .withDecodingMethod("sample")
.withTemperature(0.1f) .withTemperature(0.1)
.withTopP(0.2f) .withTopP(0.2)
.withTopK(10) .withTopK(10)
.withMaxNewTokens(30) .withMaxNewTokens(30)
.withMinNewTokens(10) .withMinNewTokens(10)
.withRepetitionPenalty(1.4f) .withRepetitionPenalty(1.4)
.withStopSequences(List.of("\n\n\n")) .withStopSequences(List.of("\n\n\n"))
.withRandomSeed(4) .withRandomSeed(4)
.build(); .build();
@@ -127,12 +127,12 @@ public class WatsonxAiChatModelTest {
WatsonxAiChatOptions modelOptions = WatsonxAiChatOptions.builder() WatsonxAiChatOptions modelOptions = WatsonxAiChatOptions.builder()
.withModel("meta-llama/llama-2-70b-chat") .withModel("meta-llama/llama-2-70b-chat")
.withDecodingMethod("sample") .withDecodingMethod("sample")
.withTemperature(0.1f) .withTemperature(0.1)
.withTopP(0.2f) .withTopP(0.2)
.withTopK(10) .withTopK(10)
.withMaxNewTokens(30) .withMaxNewTokens(30)
.withMinNewTokens(10) .withMinNewTokens(10)
.withRepetitionPenalty(1.4f) .withRepetitionPenalty(1.4)
.withStopSequences(List.of("\n\n\n")) .withStopSequences(List.of("\n\n\n"))
.withRandomSeed(4) .withRandomSeed(4)
.build(); .build();

View File

@@ -34,13 +34,13 @@ public class WatsonxAiChatOptionTest {
public void testOptions() { public void testOptions() {
WatsonxAiChatOptions options = WatsonxAiChatOptions.builder() WatsonxAiChatOptions options = WatsonxAiChatOptions.builder()
.withDecodingMethod("sample") .withDecodingMethod("sample")
.withTemperature(1.2f) .withTemperature(1.2)
.withTopK(20) .withTopK(20)
.withTopP(0.5f) .withTopP(0.5)
.withMaxNewTokens(100) .withMaxNewTokens(100)
.withMinNewTokens(20) .withMinNewTokens(20)
.withStopSequences(List.of("\n\n\n")) .withStopSequences(List.of("\n\n\n"))
.withRepetitionPenalty(1.1f) .withRepetitionPenalty(1.1)
.withRandomSeed(4) .withRandomSeed(4)
.build(); .build();
@@ -61,13 +61,13 @@ public class WatsonxAiChatOptionTest {
public void testOptionsWithAdditionalParamsOneByOne() { public void testOptionsWithAdditionalParamsOneByOne() {
WatsonxAiChatOptions options = WatsonxAiChatOptions.builder() WatsonxAiChatOptions options = WatsonxAiChatOptions.builder()
.withDecodingMethod("sample") .withDecodingMethod("sample")
.withTemperature(1.2f) .withTemperature(1.2)
.withTopK(20) .withTopK(20)
.withTopP(0.5f) .withTopP(0.5)
.withMaxNewTokens(100) .withMaxNewTokens(100)
.withMinNewTokens(20) .withMinNewTokens(20)
.withStopSequences(List.of("\n\n\n")) .withStopSequences(List.of("\n\n\n"))
.withRepetitionPenalty(1.1f) .withRepetitionPenalty(1.1)
.withRandomSeed(4) .withRandomSeed(4)
.withAdditionalProperty("HAP", true) .withAdditionalProperty("HAP", true)
.withAdditionalProperty("typicalP", 0.5f) .withAdditionalProperty("typicalP", 0.5f)
@@ -92,13 +92,13 @@ public class WatsonxAiChatOptionTest {
public void testOptionsWithAdditionalParamsMap() { public void testOptionsWithAdditionalParamsMap() {
WatsonxAiChatOptions options = WatsonxAiChatOptions.builder() WatsonxAiChatOptions options = WatsonxAiChatOptions.builder()
.withDecodingMethod("sample") .withDecodingMethod("sample")
.withTemperature(1.2f) .withTemperature(1.2)
.withTopK(20) .withTopK(20)
.withTopP(0.5f) .withTopP(0.5)
.withMaxNewTokens(100) .withMaxNewTokens(100)
.withMinNewTokens(20) .withMinNewTokens(20)
.withStopSequences(List.of("\n\n\n")) .withStopSequences(List.of("\n\n\n"))
.withRepetitionPenalty(1.1f) .withRepetitionPenalty(1.1)
.withRandomSeed(4) .withRandomSeed(4)
.withAdditionalProperties(Map.of("HAP", true, "typicalP", 0.5f, "test_value", "test")) .withAdditionalProperties(Map.of("HAP", true, "typicalP", 0.5f, "test_value", "test"))
.build(); .build();

View File

@@ -100,7 +100,7 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
*/ */
public ZhiPuAiChatModel(ZhiPuAiApi zhiPuAiApi) { public ZhiPuAiChatModel(ZhiPuAiApi zhiPuAiApi) {
this(zhiPuAiApi, this(zhiPuAiApi,
ZhiPuAiChatOptions.builder().withModel(ZhiPuAiApi.DEFAULT_CHAT_MODEL).withTemperature(0.7f).build()); ZhiPuAiChatOptions.builder().withModel(ZhiPuAiApi.DEFAULT_CHAT_MODEL).withTemperature(0.7).build());
} }
/** /**

View File

@@ -62,13 +62,13 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions, ChatOptions {
* more random, while lower values like 0.2 will make it more focused and deterministic. We generally recommend * more random, while lower values like 0.2 will make it more focused and deterministic. We generally recommend
* altering this or top_p but not both. * altering this or top_p but not both.
*/ */
private @JsonProperty("temperature") Float temperature; private @JsonProperty("temperature") Double temperature;
/** /**
* An alternative to sampling with temperature, called nucleus sampling, where the model considers the * An alternative to sampling with temperature, called nucleus sampling, where the model considers the
* results of the tokens with top_p probability mass. So 0.1 means only the tokens comprising the top 10% * results of the tokens with top_p probability mass. So 0.1 means only the tokens comprising the top 10%
* probability mass are considered. We generally recommend altering this or temperature but not both. * probability mass are considered. We generally recommend altering this or temperature but not both.
*/ */
private @JsonProperty("top_p") Float topP; private @JsonProperty("top_p") Double topP;
/** /**
* A list of tools the model may call. Currently, only functions are supported as a tool. Use this to * A list of tools the model may call. Currently, only functions are supported as a tool. Use this to
* provide a list of functions the model may generate JSON inputs for. * provide a list of functions the model may generate JSON inputs for.
@@ -156,12 +156,12 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions, ChatOptions {
return this; return this;
} }
public Builder withTemperature(Float temperature) { public Builder withTemperature(Double temperature) {
this.options.temperature = temperature; this.options.temperature = temperature;
return this; return this;
} }
public Builder withTopP(Float topP) { public Builder withTopP(Double topP) {
this.options.topP = topP; this.options.topP = topP;
return this; return this;
} }
@@ -252,20 +252,20 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions, ChatOptions {
} }
@Override @Override
public Float getTemperature() { public Double getTemperature() {
return this.temperature; return this.temperature;
} }
public void setTemperature(Float temperature) { public void setTemperature(Double temperature) {
this.temperature = temperature; this.temperature = temperature;
} }
@Override @Override
public Float getTopP() { public Double getTopP() {
return this.topP; return this.topP;
} }
public void setTopP(Float topP) { public void setTopP(Double topP) {
this.topP = topP; this.topP = topP;
} }
@@ -330,13 +330,13 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions, ChatOptions {
@Override @Override
@JsonIgnore @JsonIgnore
public Float getFrequencyPenalty() { public Double getFrequencyPenalty() {
return null; return null;
} }
@Override @Override
@JsonIgnore @JsonIgnore
public Float getPresencePenalty() { public Double getPresencePenalty() {
return null; return null;
} }

View File

@@ -46,6 +46,7 @@ import java.util.function.Predicate;
* <a href="https://open.bigmodel.cn/dev/api#text_embedding">ZhiPuAI Embedding API</a>. * <a href="https://open.bigmodel.cn/dev/api#text_embedding">ZhiPuAI Embedding API</a>.
* *
* @author Geng Rong * @author Geng Rong
* @author Thomas Vitale
* @since 1.0.0 * @since 1.0.0
*/ */
public class ZhiPuAiApi { public class ZhiPuAiApi {
@@ -230,8 +231,8 @@ public class ZhiPuAiApi {
@JsonProperty("max_tokens") Integer maxTokens, @JsonProperty("max_tokens") Integer maxTokens,
@JsonProperty("stop") List<String> stop, @JsonProperty("stop") List<String> stop,
@JsonProperty("stream") Boolean stream, @JsonProperty("stream") Boolean stream,
@JsonProperty("temperature") Float temperature, @JsonProperty("temperature") Double temperature,
@JsonProperty("top_p") Float topP, @JsonProperty("top_p") Double topP,
@JsonProperty("tools") List<FunctionTool> tools, @JsonProperty("tools") List<FunctionTool> tools,
@JsonProperty("tool_choice") Object toolChoice, @JsonProperty("tool_choice") Object toolChoice,
@JsonProperty("user") String user, @JsonProperty("user") String user,
@@ -245,7 +246,7 @@ public class ZhiPuAiApi {
* @param model ID of the model to use. * @param model ID of the model to use.
* @param temperature What sampling temperature to use, between 0 and 1. * @param temperature What sampling temperature to use, between 0 and 1.
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Float temperature) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Double temperature) {
this(messages, model, null, null, false, temperature, null, this(messages, model, null, null, false, temperature, null,
null, null, null, null, null); null, null, null, null, null);
} }
@@ -259,7 +260,7 @@ public class ZhiPuAiApi {
* @param stream If set, partial message deltas will be sent.Tokens will be sent as data-only server-sent events * @param stream If set, partial message deltas will be sent.Tokens will be sent as data-only server-sent events
* as they become available, with the stream terminated by a data: [DONE] message. * as they become available, with the stream terminated by a data: [DONE] message.
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Float temperature, boolean stream) { public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Double temperature, boolean stream) {
this(messages, model, null, null, stream, temperature, null, this(messages, model, null, null, stream, temperature, null,
null, null, null, null, null); null, null, null, null, null);
} }
@@ -275,7 +276,7 @@ public class ZhiPuAiApi {
*/ */
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model,
List<FunctionTool> tools, Object toolChoice) { List<FunctionTool> tools, Object toolChoice) {
this(messages, model, null, null, false, 0.8f, null, this(messages, model, null, null, false, 0.8, null,
tools, toolChoice, null, null, null); tools, toolChoice, null, null, null);
} }

View File

@@ -34,7 +34,7 @@ public class ChatCompletionRequestTests {
public void createRequestWithChatOptions() { public void createRequestWithChatOptions() {
var client = new ZhiPuAiChatModel(new ZhiPuAiApi("TEST"), var client = new ZhiPuAiChatModel(new ZhiPuAiApi("TEST"),
ZhiPuAiChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6f).build()); ZhiPuAiChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6).build());
var request = client.createRequest(new Prompt("Test message content"), false); var request = client.createRequest(new Prompt("Test message content"), false);
@@ -42,16 +42,16 @@ public class ChatCompletionRequestTests {
assertThat(request.stream()).isFalse(); assertThat(request.stream()).isFalse();
assertThat(request.model()).isEqualTo("DEFAULT_MODEL"); assertThat(request.model()).isEqualTo("DEFAULT_MODEL");
assertThat(request.temperature()).isEqualTo(66.6f); assertThat(request.temperature()).isEqualTo(66.6);
request = client.createRequest(new Prompt("Test message content", request = client.createRequest(new Prompt("Test message content",
ZhiPuAiChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9f).build()), true); ZhiPuAiChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9).build()), true);
assertThat(request.messages()).hasSize(1); assertThat(request.messages()).hasSize(1);
assertThat(request.stream()).isTrue(); assertThat(request.stream()).isTrue();
assertThat(request.model()).isEqualTo("PROMPT_MODEL"); assertThat(request.model()).isEqualTo("PROMPT_MODEL");
assertThat(request.temperature()).isEqualTo(99.9f); assertThat(request.temperature()).isEqualTo(99.9);
} }
@Test @Test

Some files were not shown because too many files have changed in this diff Show More