From bc9773ff9aab32c06220204757766aa2150d302a Mon Sep 17 00:00:00 2001 From: Roo Code Date: Fri, 7 Feb 2025 23:52:30 -0500 Subject: [PATCH] Configure per-configuration temperature --- src/api/providers/anthropic.ts | 6 +- src/api/providers/bedrock.ts | 4 +- src/api/providers/gemini.ts | 4 +- src/api/providers/glama.ts | 4 +- src/api/providers/lmstudio.ts | 4 +- src/api/providers/mistral.ts | 2 +- src/api/providers/ollama.ts | 12 +-- src/api/providers/openai-native.ts | 4 +- src/api/providers/openai.ts | 2 +- src/api/providers/openrouter.ts | 11 ++- src/api/providers/unbound.ts | 4 +- src/api/providers/vertex.ts | 4 +- src/core/webview/ClineProvider.ts | 6 ++ src/shared/api.ts | 1 + .../src/components/settings/ApiOptions.tsx | 13 ++++ .../settings/TemperatureControl.tsx | 73 +++++++++++++++++++ 16 files changed, 124 insertions(+), 30 deletions(-) create mode 100644 webview-ui/src/components/settings/TemperatureControl.tsx diff --git a/src/api/providers/anthropic.ts b/src/api/providers/anthropic.ts index e65b82ddef..2059804a54 100644 --- a/src/api/providers/anthropic.ts +++ b/src/api/providers/anthropic.ts @@ -44,7 +44,7 @@ export class AnthropicHandler implements ApiHandler, SingleCompletionHandler { { model: modelId, max_tokens: this.getModel().info.maxTokens || 8192, - temperature: 0, + temperature: this.options.modelTemperature ?? 0, system: [{ text: systemPrompt, type: "text", cache_control: { type: "ephemeral" } }], // setting cache breakpoint for system prompt so new tasks can reuse it messages: messages.map((message, index) => { if (index === lastUserMsgIndex || index === secondLastMsgUserIndex) { @@ -96,7 +96,7 @@ export class AnthropicHandler implements ApiHandler, SingleCompletionHandler { stream = (await this.client.messages.create({ model: modelId, max_tokens: this.getModel().info.maxTokens || 8192, - temperature: 0, + temperature: this.options.modelTemperature ?? 0, system: [{ text: systemPrompt, type: "text" }], messages, // tools, @@ -179,7 +179,7 @@ export class AnthropicHandler implements ApiHandler, SingleCompletionHandler { const response = await this.client.messages.create({ model: this.getModel().id, max_tokens: this.getModel().info.maxTokens || 8192, - temperature: 0, + temperature: this.options.modelTemperature ?? 0, messages: [{ role: "user", content: prompt }], stream: false, }) diff --git a/src/api/providers/bedrock.ts b/src/api/providers/bedrock.ts index 0e90c2bcc4..17362e1f05 100644 --- a/src/api/providers/bedrock.ts +++ b/src/api/providers/bedrock.ts @@ -104,7 +104,7 @@ export class AwsBedrockHandler implements ApiHandler, SingleCompletionHandler { system: [{ text: systemPrompt }], inferenceConfig: { maxTokens: modelConfig.info.maxTokens || 5000, - temperature: 0.3, + temperature: this.options.modelTemperature ?? 0.3, topP: 0.1, ...(this.options.awsUsePromptCache ? { @@ -262,7 +262,7 @@ export class AwsBedrockHandler implements ApiHandler, SingleCompletionHandler { ]), inferenceConfig: { maxTokens: modelConfig.info.maxTokens || 5000, - temperature: 0.3, + temperature: this.options.modelTemperature ?? 0.3, topP: 0.1, }, } diff --git a/src/api/providers/gemini.ts b/src/api/providers/gemini.ts index 0577a021e6..e9a0015224 100644 --- a/src/api/providers/gemini.ts +++ b/src/api/providers/gemini.ts @@ -23,7 +23,7 @@ export class GeminiHandler implements ApiHandler, SingleCompletionHandler { contents: messages.map(convertAnthropicMessageToGemini), generationConfig: { // maxOutputTokens: this.getModel().info.maxTokens, - temperature: 0, + temperature: this.options.modelTemperature ?? 0, }, }) @@ -60,7 +60,7 @@ export class GeminiHandler implements ApiHandler, SingleCompletionHandler { const result = await model.generateContent({ contents: [{ role: "user", parts: [{ text: prompt }] }], generationConfig: { - temperature: 0, + temperature: this.options.modelTemperature ?? 0, }, }) diff --git a/src/api/providers/glama.ts b/src/api/providers/glama.ts index 95b806f27c..226891b16a 100644 --- a/src/api/providers/glama.ts +++ b/src/api/providers/glama.ts @@ -79,7 +79,7 @@ export class GlamaHandler implements ApiHandler, SingleCompletionHandler { } if (this.supportsTemperature()) { - requestOptions.temperature = 0 + requestOptions.temperature = this.options.modelTemperature ?? 0 } const { data: completion, response } = await this.client.chat.completions @@ -172,7 +172,7 @@ export class GlamaHandler implements ApiHandler, SingleCompletionHandler { } if (this.supportsTemperature()) { - requestOptions.temperature = 0 + requestOptions.temperature = this.options.modelTemperature ?? 0 } if (this.getModel().id.startsWith("anthropic/")) { diff --git a/src/api/providers/lmstudio.ts b/src/api/providers/lmstudio.ts index 81cec81b4d..cc164d240e 100644 --- a/src/api/providers/lmstudio.ts +++ b/src/api/providers/lmstudio.ts @@ -27,7 +27,7 @@ export class LmStudioHandler implements ApiHandler, SingleCompletionHandler { const stream = await this.client.chat.completions.create({ model: this.getModel().id, messages: openAiMessages, - temperature: 0, + temperature: this.options.modelTemperature ?? 0, stream: true, }) for await (const chunk of stream) { @@ -59,7 +59,7 @@ export class LmStudioHandler implements ApiHandler, SingleCompletionHandler { const response = await this.client.chat.completions.create({ model: this.getModel().id, messages: [{ role: "user", content: prompt }], - temperature: 0, + temperature: this.options.modelTemperature ?? 0, stream: false, }) return response.choices[0]?.message.content || "" diff --git a/src/api/providers/mistral.ts b/src/api/providers/mistral.ts index c4377f0003..4bcf1a191c 100644 --- a/src/api/providers/mistral.ts +++ b/src/api/providers/mistral.ts @@ -30,7 +30,7 @@ export class MistralHandler implements ApiHandler { const stream = await this.client.chat.stream({ model: this.getModel().id, // max_completion_tokens: this.getModel().info.maxTokens, - temperature: 0, + temperature: this.options.modelTemperature ?? 0, messages: [{ role: "system", content: systemPrompt }, ...convertToMistralMessages(messages)], stream: true, }) diff --git a/src/api/providers/ollama.ts b/src/api/providers/ollama.ts index 4175b78fa5..44e0d3f4af 100644 --- a/src/api/providers/ollama.ts +++ b/src/api/providers/ollama.ts @@ -20,7 +20,7 @@ export class OllamaHandler implements ApiHandler, SingleCompletionHandler { async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream { const modelId = this.getModel().id - const useR1Format = modelId.toLowerCase().includes('deepseek-r1') + const useR1Format = modelId.toLowerCase().includes("deepseek-r1") const openAiMessages: OpenAI.Chat.ChatCompletionMessageParam[] = [ { role: "system", content: systemPrompt }, ...(useR1Format ? convertToR1Format(messages) : convertToOpenAiMessages(messages)), @@ -29,7 +29,7 @@ export class OllamaHandler implements ApiHandler, SingleCompletionHandler { const stream = await this.client.chat.completions.create({ model: this.getModel().id, messages: openAiMessages, - temperature: 0, + temperature: this.options.modelTemperature ?? 0, stream: true, }) for await (const chunk of stream) { @@ -53,11 +53,13 @@ export class OllamaHandler implements ApiHandler, SingleCompletionHandler { async completePrompt(prompt: string): Promise { try { const modelId = this.getModel().id - const useR1Format = modelId.toLowerCase().includes('deepseek-r1') + const useR1Format = modelId.toLowerCase().includes("deepseek-r1") const response = await this.client.chat.completions.create({ model: this.getModel().id, - messages: useR1Format ? convertToR1Format([{ role: "user", content: prompt }]) : [{ role: "user", content: prompt }], - temperature: 0, + messages: useR1Format + ? convertToR1Format([{ role: "user", content: prompt }]) + : [{ role: "user", content: prompt }], + temperature: this.options.modelTemperature ?? 0, stream: false, }) return response.choices[0]?.message.content || "" diff --git a/src/api/providers/openai-native.ts b/src/api/providers/openai-native.ts index e4883b7a98..a40e002ce1 100644 --- a/src/api/providers/openai-native.ts +++ b/src/api/providers/openai-native.ts @@ -88,7 +88,7 @@ export class OpenAiNativeHandler implements ApiHandler, SingleCompletionHandler ): ApiStream { const stream = await this.client.chat.completions.create({ model: modelId, - temperature: 0, + temperature: this.options.modelTemperature ?? 0, messages: [{ role: "system", content: systemPrompt }, ...convertToOpenAiMessages(messages)], stream: true, stream_options: { include_usage: true }, @@ -189,7 +189,7 @@ export class OpenAiNativeHandler implements ApiHandler, SingleCompletionHandler return { model: modelId, messages: [{ role: "user", content: prompt }], - temperature: 0, + temperature: this.options.modelTemperature ?? 0, } } } diff --git a/src/api/providers/openai.ts b/src/api/providers/openai.ts index 408f4e5cc3..da3cf1b9e7 100644 --- a/src/api/providers/openai.ts +++ b/src/api/providers/openai.ts @@ -57,7 +57,7 @@ export class OpenAiHandler implements ApiHandler, SingleCompletionHandler { } const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming = { model: modelId, - temperature: 0, + temperature: this.options.modelTemperature ?? (deepseekReasoner ? 0.6 : 0), messages: deepseekReasoner ? convertToR1Format([{ role: "user", content: systemPrompt }, ...messages]) : [systemMessage, ...convertToOpenAiMessages(messages)], diff --git a/src/api/providers/openrouter.ts b/src/api/providers/openrouter.ts index 0e23c5d35d..fa1c65d126 100644 --- a/src/api/providers/openrouter.ts +++ b/src/api/providers/openrouter.ts @@ -115,7 +115,7 @@ export class OpenRouterHandler implements ApiHandler, SingleCompletionHandler { break } - let temperature = 0 + let defaultTemperature = 0 let topP: number | undefined = undefined // Handle models based on deepseek-r1 @@ -124,9 +124,8 @@ export class OpenRouterHandler implements ApiHandler, SingleCompletionHandler { this.getModel().id === "perplexity/sonar-reasoning" ) { // Recommended temperature for DeepSeek reasoning models - temperature = 0.6 - // DeepSeek highly recommends using user instead of system - // role + defaultTemperature = 0.6 + // DeepSeek highly recommends using user instead of system role openAiMessages = convertToR1Format([{ role: "user", content: systemPrompt }, ...messages]) // Some provider support topP and 0.95 is value that Deepseek used in their benchmarks topP = 0.95 @@ -137,7 +136,7 @@ export class OpenRouterHandler implements ApiHandler, SingleCompletionHandler { const stream = await this.client.chat.completions.create({ model: this.getModel().id, max_tokens: maxTokens, - temperature: temperature, + temperature: this.options.modelTemperature ?? defaultTemperature, top_p: topP, messages: openAiMessages, stream: true, @@ -224,7 +223,7 @@ export class OpenRouterHandler implements ApiHandler, SingleCompletionHandler { const response = await this.client.chat.completions.create({ model: this.getModel().id, messages: [{ role: "user", content: prompt }], - temperature: 0, + temperature: this.options.modelTemperature ?? 0, stream: false, }) diff --git a/src/api/providers/unbound.ts b/src/api/providers/unbound.ts index 305bd282ad..809ae21889 100644 --- a/src/api/providers/unbound.ts +++ b/src/api/providers/unbound.ts @@ -79,7 +79,7 @@ export class UnboundHandler implements ApiHandler, SingleCompletionHandler { { model: this.getModel().id.split("/")[1], max_tokens: maxTokens, - temperature: 0, + temperature: this.options.modelTemperature ?? 0, messages: openAiMessages, stream: true, }, @@ -146,7 +146,7 @@ export class UnboundHandler implements ApiHandler, SingleCompletionHandler { const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsNonStreaming = { model: this.getModel().id.split("/")[1], messages: [{ role: "user", content: prompt }], - temperature: 0, + temperature: this.options.modelTemperature ?? 0, } if (this.getModel().id.startsWith("anthropic/")) { diff --git a/src/api/providers/vertex.ts b/src/api/providers/vertex.ts index 1ea68eaa4e..0ee22e5893 100644 --- a/src/api/providers/vertex.ts +++ b/src/api/providers/vertex.ts @@ -22,7 +22,7 @@ export class VertexHandler implements ApiHandler, SingleCompletionHandler { const stream = await this.client.messages.create({ model: this.getModel().id, max_tokens: this.getModel().info.maxTokens || 8192, - temperature: 0, + temperature: this.options.modelTemperature ?? 0, system: systemPrompt, messages, stream: true, @@ -89,7 +89,7 @@ export class VertexHandler implements ApiHandler, SingleCompletionHandler { const response = await this.client.messages.create({ model: this.getModel().id, max_tokens: this.getModel().info.maxTokens || 8192, - temperature: 0, + temperature: this.options.modelTemperature ?? 0, messages: [{ role: "user", content: prompt }], stream: false, }) diff --git a/src/core/webview/ClineProvider.ts b/src/core/webview/ClineProvider.ts index 23346d945c..221c89a9ad 100644 --- a/src/core/webview/ClineProvider.ts +++ b/src/core/webview/ClineProvider.ts @@ -119,6 +119,7 @@ type GlobalStateKey = | "autoApprovalEnabled" | "customModes" // Array of custom modes | "unboundModelId" + | "modelTemperature" export const GlobalFileNames = { apiConversationHistory: "api_conversation_history.json", @@ -1538,6 +1539,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { mistralApiKey, unboundApiKey, unboundModelId, + modelTemperature, } = apiConfiguration await this.updateGlobalState("apiProvider", apiProvider) await this.updateGlobalState("apiModelId", apiModelId) @@ -1578,6 +1580,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { await this.storeSecret("mistralApiKey", mistralApiKey) await this.storeSecret("unboundApiKey", unboundApiKey) await this.updateGlobalState("unboundModelId", unboundModelId) + await this.updateGlobalState("modelTemperature", modelTemperature) if (this.cline) { this.cline.api = buildApiHandler(apiConfiguration) } @@ -2254,6 +2257,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { experiments, unboundApiKey, unboundModelId, + modelTemperature, ] = await Promise.all([ this.getGlobalState("apiProvider") as Promise, this.getGlobalState("apiModelId") as Promise, @@ -2328,6 +2332,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { this.getGlobalState("experiments") as Promise | undefined>, this.getSecret("unboundApiKey") as Promise, this.getGlobalState("unboundModelId") as Promise, + this.getGlobalState("modelTemperature") as Promise, ]) let apiProvider: ApiProvider @@ -2385,6 +2390,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { vsCodeLmModelSelector, unboundApiKey, unboundModelId, + modelTemperature, }, lastShownAnnouncementId, customInstructions, diff --git a/src/shared/api.ts b/src/shared/api.ts index 39bc2a69ca..c4f62259a7 100644 --- a/src/shared/api.ts +++ b/src/shared/api.ts @@ -60,6 +60,7 @@ export interface ApiHandlerOptions { includeMaxTokens?: boolean unboundApiKey?: string unboundModelId?: string + modelTemperature?: number } export type ApiConfiguration = ApiHandlerOptions & { diff --git a/webview-ui/src/components/settings/ApiOptions.tsx b/webview-ui/src/components/settings/ApiOptions.tsx index 4f9c8e1b78..7e51347c0b 100644 --- a/webview-ui/src/components/settings/ApiOptions.tsx +++ b/webview-ui/src/components/settings/ApiOptions.tsx @@ -2,6 +2,7 @@ import { memo, useCallback, useEffect, useMemo, useState } from "react" import { useEvent, useInterval } from "react-use" import { Checkbox, Dropdown, Pane, type DropdownOption } from "vscrui" import { VSCodeLink, VSCodeRadio, VSCodeRadioGroup, VSCodeTextField } from "@vscode/webview-ui-toolkit/react" +import { TemperatureControl } from "./TemperatureControl" import * as vscodemodels from "vscode" import { @@ -1361,6 +1362,18 @@ const ApiOptions = ({ apiErrorMessage, modelIdErrorMessage }: ApiOptionsProps) = )} +
+ { + handleInputChange("modelTemperature")({ + target: { value }, + }) + }} + maxValue={2} + /> +
+ {modelIdErrorMessage && (

void + maxValue?: number // Some providers like OpenAI use 0-2 range +} + +export const TemperatureControl = ({ value, onChange, maxValue = 1 }: TemperatureControlProps) => { + const [isCustomTemperature, setIsCustomTemperature] = useState(value !== undefined) + + // Sync internal state with prop changes when switching profiles + useEffect(() => { + setIsCustomTemperature(value !== undefined) + }, [value]) + + return ( +

+ { + setIsCustomTemperature(checked) + if (!checked) { + onChange(undefined) // Reset to provider default + } else { + onChange(0) // Set initial value when enabling + } + }}> + Use custom temperature + + + {isCustomTemperature && ( + <> + + { + const newValue = parseFloat(e.target.value) + onChange(isNaN(newValue) ? undefined : newValue) + }} + style={{ + flexGrow: 1, + accentColor: "var(--vscode-button-background)", + height: "2px", + }} + /> + + {value?.toFixed(2)} + + + )} +
+ ) +}