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

@@ -100,7 +100,7 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
*/
public ZhiPuAiChatModel(ZhiPuAiApi 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
* 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
* 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.
*/
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
* provide a list of functions the model may generate JSON inputs for.
@@ -156,12 +156,12 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions, ChatOptions {
return this;
}
public Builder withTemperature(Float temperature) {
public Builder withTemperature(Double temperature) {
this.options.temperature = temperature;
return this;
}
public Builder withTopP(Float topP) {
public Builder withTopP(Double topP) {
this.options.topP = topP;
return this;
}
@@ -252,20 +252,20 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions, ChatOptions {
}
@Override
public Float getTemperature() {
public Double getTemperature() {
return this.temperature;
}
public void setTemperature(Float temperature) {
public void setTemperature(Double temperature) {
this.temperature = temperature;
}
@Override
public Float getTopP() {
public Double getTopP() {
return this.topP;
}
public void setTopP(Float topP) {
public void setTopP(Double topP) {
this.topP = topP;
}
@@ -330,13 +330,13 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions, ChatOptions {
@Override
@JsonIgnore
public Float getFrequencyPenalty() {
public Double getFrequencyPenalty() {
return null;
}
@Override
@JsonIgnore
public Float getPresencePenalty() {
public Double getPresencePenalty() {
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>.
*
* @author Geng Rong
* @author Thomas Vitale
* @since 1.0.0
*/
public class ZhiPuAiApi {
@@ -230,8 +231,8 @@ public class ZhiPuAiApi {
@JsonProperty("max_tokens") Integer maxTokens,
@JsonProperty("stop") List<String> stop,
@JsonProperty("stream") Boolean stream,
@JsonProperty("temperature") Float temperature,
@JsonProperty("top_p") Float topP,
@JsonProperty("temperature") Double temperature,
@JsonProperty("top_p") Double topP,
@JsonProperty("tools") List<FunctionTool> tools,
@JsonProperty("tool_choice") Object toolChoice,
@JsonProperty("user") String user,
@@ -245,7 +246,7 @@ public class ZhiPuAiApi {
* @param model ID of the model to use.
* @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,
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
* 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,
null, null, null, null, null);
}
@@ -275,7 +276,7 @@ public class ZhiPuAiApi {
*/
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model,
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);
}

View File

@@ -34,7 +34,7 @@ public class ChatCompletionRequestTests {
public void createRequestWithChatOptions() {
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);
@@ -42,16 +42,16 @@ public class ChatCompletionRequestTests {
assertThat(request.stream()).isFalse();
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",
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.stream()).isTrue();
assertThat(request.model()).isEqualTo("PROMPT_MODEL");
assertThat(request.temperature()).isEqualTo(99.9f);
assertThat(request.temperature()).isEqualTo(99.9);
}
@Test

View File

@@ -43,8 +43,8 @@ public class ZhiPuAiApiIT {
@Test
void chatCompletionEntity() {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER);
ResponseEntity<ChatCompletion> response = zhiPuAiApi.chatCompletionEntity(
new ChatCompletionRequest(List.of(chatCompletionMessage), "glm-3-turbo", 0.7f, false));
ResponseEntity<ChatCompletion> response = zhiPuAiApi
.chatCompletionEntity(new ChatCompletionRequest(List.of(chatCompletionMessage), "glm-3-turbo", 0.7, false));
assertThat(response).isNotNull();
assertThat(response.getBody()).isNotNull();
@@ -55,7 +55,7 @@ public class ZhiPuAiApiIT {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER);
ResponseEntity<ChatCompletion> response = zhiPuAiApi
.chatCompletionEntity(new ChatCompletionRequest(List.of(chatCompletionMessage), "glm-3-turbo", 1024, null,
false, 0.95f, 0.7f, null, null, null, "test_request_id", false));
false, 0.95, 0.7, null, null, null, "test_request_id", false));
assertThat(response).isNotNull();
assertThat(response.getBody()).isNotNull();
@@ -65,7 +65,7 @@ public class ZhiPuAiApiIT {
void chatCompletionStream() {
ChatCompletionMessage chatCompletionMessage = new ChatCompletionMessage("Hello world", Role.USER);
Flux<ChatCompletionChunk> response = zhiPuAiApi
.chatCompletionStream(new ChatCompletionRequest(List.of(chatCompletionMessage), "glm-3-turbo", 0.7f, true));
.chatCompletionStream(new ChatCompletionRequest(List.of(chatCompletionMessage), "glm-3-turbo", 0.7, true));
assertThat(response).isNotNull();
assertThat(response.collectList().block()).isNotNull();