other things + Embedding Model selection

Alessandro committed Nov 4, 2024 at 22:26 UTC 61b5b8389a8af536c0444b4a659be9d04c1cff06
4 files changed +87 -14
initialize.py
+1 -1
@@ -33,7 +33,7 @@ def initialize():
33 chat_model = chat_llm,
34 utility_model = utility_llm,
35 embeddings_model = embedding_llm,
36 - prompts_subdir = "dianoia-xl",
36 + prompts_subdir = "",
37 # memory_subdir = "",
38 knowledge_subdirs = ["default","custom"],
39 auto_memory_count = 0,
models.py
+15
@@ -44,6 +44,13 @@ class ModelProvider(Enum):
44 OPENROUTER = "OpenRouter"
45 SAMBANOVA = "Sambanova"
46
47 +class EmbeddingProvider(Enum):
48 + OPENAI = "OpenAI" # default
49 + HUGGINGFACE = "HuggingFace"
50 + OLLAMA = "Ollama"
51 + LMSTUDIO = "LM Studio"
52 + OPENROUTER = "OpenRouter"
53 + AZURE = "OpenAI Azure"
54
55 # Utility function to get API keys from environment variables
56 def get_api_key(service):
@@ -53,10 +60,18 @@ def get_api_key(service):
60
61
62 def get_model(type: ModelType, provider: ModelProvider, name: str, **kwargs):
63 + if type == ModelType.EMBEDDING:
64 + # call function for embedding models
65 + return get_embedding_model(provider, name, **kwargs)
66 + # for other model types
67 fnc_name = f"get_{provider.name.lower()}_{type.name.lower()}" # function name of model getter
68 model = globals()[fnc_name](name, **kwargs) # call function by name
69 return model
70
71 +def get_embedding_model(provider: EmbeddingProvider, name: str, **kwargs):
72 + fnc_name = f"get_{provider.name.lower()}_embedding" # function name for embedding models
73 + model = globals()[fnc_name](name, **kwargs) # call function by name
74 + return model
75
76 # Ollama models
77 def get_ollama_chat(
python/helpers/settings.py
+46 -9
@@ -3,7 +3,7 @@ import os
3 import re
4 from typing import Any, Optional, TypedDict
5 from . import files
6 -from models import get_model, ModelProvider, ModelType
6 +from models import get_model, get_embedding_model, ModelProvider, EmbeddingProvider, ModelType
7 from langchain_core.language_models.chat_models import BaseChatModel
8 from langchain_core.embeddings import Embeddings
9
@@ -32,7 +32,6 @@ _settings: Settings | None = None
32
33
34 def convert_out(settings: Settings) -> dict[str, Any]:
35 -
35 # main model section
36 chat_model_fields = []
37 chat_model_fields.append(
@@ -78,13 +77,13 @@ def convert_out(settings: Settings) -> dict[str, Any]:
77 }
78 )
79
81 - chat_model_seciton = {
80 + chat_model_section = {
81 "title": "Chat Model",
82 "description": "Selection and settings for main chat model used by Agent Zero",
83 "fields": chat_model_fields,
84 }
85
87 - # main model section
86 + # utility model section
87 util_model_fields = []
88 util_model_fields.append(
89 {
@@ -129,13 +128,51 @@ def convert_out(settings: Settings) -> dict[str, Any]:
128 }
129 )
130
132 - util_model_seciton = {
133 - "title": "Utility model",
131 + util_model_section = {
132 + "title": "Utility Model",
133 "description": "Smaller, cheaper, faster model for handling utility tasks like organizing memory, preparing prompts, summarizing.",
134 "fields": util_model_fields,
135 }
136
138 - result = {"sections": [chat_model_seciton, util_model_seciton]}
137 + # embedding model section
138 + embed_model_fields = []
139 + embed_model_fields.append(
140 + {
141 + "id": "embed_model_provider",
142 + "title": "Embedding model provider",
143 + "description": "Select provider for embedding model used by the framework",
144 + "type": "select",
145 + "value": settings["embed_model_provider"],
146 + "options": [{"value": p.name, "label": p.value} for p in EmbeddingProvider],
147 + }
148 + )
149 + embed_model_fields.append(
150 + {
151 + "id": "embed_model_name",
152 + "title": "Embedding model name",
153 + "description": "Exact name of model from selected provider",
154 + "type": "input",
155 + "value": settings["embed_model_name"],
156 + }
157 + )
158 +
159 + embed_model_fields.append(
160 + {
161 + "id": "embed_model_kwargs",
162 + "title": "Embedding model additional parameters",
163 + "description": "Any other parameters supported by the model. Format is KEY=VALUE on individual lines, just like .env file.",
164 + "type": "textarea",
165 + "value": _dict_to_env(settings["embed_model_kwargs"]),
166 + }
167 + )
168 +
169 + embed_model_section = {
170 + "title": "Embedding Model",
171 + "description": "Settings for the embedding model used by Agent Zero.",
172 + "fields": embed_model_fields,
173 + }
174 +
175 + result = {"sections": [chat_model_section, util_model_section, embed_model_section]}
176 return result
177
178 def convert_in(settings: dict[str, Any]) -> Settings:
@@ -200,7 +237,7 @@ def get_embedding_model() -> Embeddings:
237 settings = get_settings()
238 return get_model(
239 type=ModelType.EMBEDDING,
203 - provider=ModelProvider[settings["embed_model_provider"]],
240 + provider=EmbeddingProvider[settings["embed_model_provider"]],
241 name=settings["embed_model_name"],
242 **settings["embed_model_kwargs"],
243 )
@@ -228,7 +265,7 @@ def _get_default_settings() -> Settings:
265 util_model_name="gpt-4o-mini",
266 util_model_temperature=0,
267 util_model_kwargs={},
231 - embed_model_provider=ModelProvider.OPENAI.name,
268 + embed_model_provider=EmbeddingProvider.OPENAI.name,
269 embed_model_name="text-embedding-3-small",
270 embed_model_kwargs={},
271 )
webui/settings.css
+25 -4
@@ -1,3 +1,11 @@
1 +* {
2 + transition: all var(--transition-speed) ease-in-out;
3 +}
4 +
5 +select {
6 + transition: none;
7 +}
8 +
9 .modal-overlay {
10 position: fixed;
11 top: 0;
@@ -23,8 +31,8 @@
31 }
32
33 .modal-header {
26 - padding: 1.5rem 2rem;
27 - border-bottom: 1px solid #eee;
34 + padding: 0.875rem 2rem;
35 + border-bottom: 1px solid var(--color-border);
36 }
37
38 .modal-content {
@@ -34,6 +42,8 @@
42 background-clip: border-box;
43 border: 6px solid transparent;
44 transition: all 0.3s ease;
45 + margin-bottom: 0;
46 + padding-bottom: 0;
47 }
48
49 .modal-content::-webkit-scrollbar {
@@ -212,19 +222,29 @@ input[type="range"] {
222 .btn-ok {
223 background: #3270e2;
224 color: white;
225 + transition: background 0.3s ease-in-out;
226 }
227
228 .btn-ok:hover{
218 - background: #274170;
229 + background: #3265c0;
230 +}
231 +
232 +.btn-ok:active{
233 + background: #345693;
234 }
235
236 .btn-cancel {
237 background: #ddd;
238 color: #333;
239 + transition: background 0.3s ease-in-out;
240 }
241
242 .btn-cancel:hover {
227 - background: #222
243 + background: #a6a6a6
244 +}
245 +
246 +.btn-cancel:active {
247 + background: #808080
248 }
249
250 .btn-field {
@@ -247,6 +267,7 @@ select {
267 font-size: inherit;
268 cursor: pointer;
269 font-family: Rubik;
270 + outline: none;
271 }
272
273 select:disabled {