revised models page
This commit is contained in:
parent
6532411d33
commit
e4e2384587
1 changed files with 211 additions and 200 deletions
|
|
@ -1,10 +1,10 @@
|
||||||
import os
|
import os
|
||||||
|
|
||||||
import streamlit as st
|
import streamlit as st
|
||||||
|
from esperanto import AIFactory
|
||||||
|
|
||||||
from open_notebook.config import CONFIG
|
from open_notebook.config import CONFIG
|
||||||
from open_notebook.domain.models import DefaultModels, Model, model_manager
|
from open_notebook.domain.models import DefaultModels, Model, model_manager
|
||||||
from open_notebook.models import MODEL_CLASS_MAP
|
|
||||||
from pages.components.model_selector import model_selector
|
from pages.components.model_selector import model_selector
|
||||||
from pages.stream_app.utils import setup_page
|
from pages.stream_app.utils import setup_page
|
||||||
|
|
||||||
|
|
@ -13,8 +13,6 @@ setup_page("🤖 Models", only_check_mandatory_models=False, stop_on_model_error
|
||||||
|
|
||||||
st.title("🤖 Models")
|
st.title("🤖 Models")
|
||||||
|
|
||||||
model_tab, model_defaults_tab = st.tabs(["Models", "Model Defaults"])
|
|
||||||
|
|
||||||
provider_status = {}
|
provider_status = {}
|
||||||
|
|
||||||
model_types = [
|
model_types = [
|
||||||
|
|
@ -25,6 +23,8 @@ model_types = [
|
||||||
"speech_to_text",
|
"speech_to_text",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def check_available_providers():
|
||||||
provider_status["ollama"] = os.environ.get("OLLAMA_API_BASE") is not None
|
provider_status["ollama"] = os.environ.get("OLLAMA_API_BASE") is not None
|
||||||
provider_status["openai"] = os.environ.get("OPENAI_API_KEY") is not None
|
provider_status["openai"] = os.environ.get("OPENAI_API_KEY") is not None
|
||||||
provider_status["groq"] = os.environ.get("GROQ_API_KEY") is not None
|
provider_status["groq"] = os.environ.get("GROQ_API_KEY") is not None
|
||||||
|
|
@ -34,11 +34,11 @@ provider_status["vertexai"] = (
|
||||||
and os.environ.get("VERTEX_LOCATION") is not None
|
and os.environ.get("VERTEX_LOCATION") is not None
|
||||||
and os.environ.get("GOOGLE_APPLICATION_CREDENTIALS") is not None
|
and os.environ.get("GOOGLE_APPLICATION_CREDENTIALS") is not None
|
||||||
)
|
)
|
||||||
provider_status["vertexai-anthropic"] = (
|
# provider_status["vertexai-anthropic"] = (
|
||||||
os.environ.get("VERTEX_PROJECT") is not None
|
# os.environ.get("VERTEX_PROJECT") is not None
|
||||||
and os.environ.get("VERTEX_LOCATION") is not None
|
# and os.environ.get("VERTEX_LOCATION") is not None
|
||||||
and os.environ.get("GOOGLE_APPLICATION_CREDENTIALS") is not None
|
# and os.environ.get("GOOGLE_APPLICATION_CREDENTIALS") is not None
|
||||||
)
|
# )
|
||||||
provider_status["gemini"] = os.environ.get("GOOGLE_API_KEY") is not None
|
provider_status["gemini"] = os.environ.get("GOOGLE_API_KEY") is not None
|
||||||
provider_status["openrouter"] = (
|
provider_status["openrouter"] = (
|
||||||
os.environ.get("OPENROUTER_API_KEY") is not None
|
os.environ.get("OPENROUTER_API_KEY") is not None
|
||||||
|
|
@ -47,18 +47,21 @@ provider_status["openrouter"] = (
|
||||||
)
|
)
|
||||||
provider_status["anthropic"] = os.environ.get("ANTHROPIC_API_KEY") is not None
|
provider_status["anthropic"] = os.environ.get("ANTHROPIC_API_KEY") is not None
|
||||||
provider_status["elevenlabs"] = os.environ.get("ELEVENLABS_API_KEY") is not None
|
provider_status["elevenlabs"] = os.environ.get("ELEVENLABS_API_KEY") is not None
|
||||||
provider_status["litellm"] = (
|
provider_status["voyage"] = os.environ.get("VORAGE_API_KEY") is not None
|
||||||
provider_status["ollama"]
|
provider_status["azure"] = (
|
||||||
or provider_status["vertexai"]
|
os.environ.get("AZURE_OPENAI_API_KEY") is not None
|
||||||
or provider_status["vertexai-anthropic"]
|
and os.environ.get("AZURE_OPENAI_ENDPOINT") is not None
|
||||||
or provider_status["anthropic"]
|
and os.environ.get("AZURE_OPENAI_DEPLOYMENT_NAME") is not None
|
||||||
or provider_status["openai"]
|
and os.environ.get("AZURE_OPENAI_API_VERSION") is not None
|
||||||
or provider_status["gemini"]
|
|
||||||
)
|
)
|
||||||
|
provider_status["mistral"] = os.environ.get("MISTRAL_API_KEY") is not None
|
||||||
|
provider_status["deepseek"] = os.environ.get("DEEPSEEK_API_KEY") is not None
|
||||||
|
|
||||||
available_providers = [k for k, v in provider_status.items() if v]
|
available_providers = [k for k, v in provider_status.items() if v]
|
||||||
unavailable_providers = [k for k, v in provider_status.items() if not v]
|
unavailable_providers = [k for k, v in provider_status.items() if not v]
|
||||||
|
|
||||||
|
return available_providers, unavailable_providers
|
||||||
|
|
||||||
|
|
||||||
def generate_new_models(models, suggested_models):
|
def generate_new_models(models, suggested_models):
|
||||||
# Create a set of existing model keys for efficient lookup
|
# Create a set of existing model keys for efficient lookup
|
||||||
|
|
@ -91,30 +94,33 @@ def generate_new_models(models, suggested_models):
|
||||||
|
|
||||||
default_models = DefaultModels()
|
default_models = DefaultModels()
|
||||||
all_models = Model.get_all()
|
all_models = Model.get_all()
|
||||||
|
esperanto_available_providers = AIFactory.get_available_providers()
|
||||||
|
|
||||||
with model_tab:
|
|
||||||
|
st.subheader("Provider Availability")
|
||||||
|
st.markdown(
|
||||||
|
"Below, you'll find all AI providers supported and their current availability status. To enable more providers, you need to setup some of their ENV Variables. Please check [the documentation](https://github.com/lfnovo/open-notebook) for instructions on how to do so."
|
||||||
|
)
|
||||||
|
available_providers, unavailable_providers = check_available_providers()
|
||||||
|
with st.expander("Available Providers"):
|
||||||
|
st.write(available_providers)
|
||||||
|
with st.expander("Unavailable Providers"):
|
||||||
|
st.write(unavailable_providers)
|
||||||
|
|
||||||
|
st.divider()
|
||||||
st.subheader("Add Model")
|
st.subheader("Add Model")
|
||||||
|
st.markdown(
|
||||||
provider = st.selectbox("Provider", available_providers)
|
"Even though a lot of models can be supported, not all will perform optimally. Some are more fit for use in this tool than others. To help you decide which models to use, please refer to [Which model to choose?](https://github.com/lfnovo/open-notebook/blob/main/docs/SETUP.md#which-model-to-choose) for more information. You can also play with some models in the [Transformations](https://try-it-out.open-notebook.com) page to see if they match your needs."
|
||||||
if len(unavailable_providers) > 0:
|
|
||||||
st.caption(
|
|
||||||
f"Unavailable Providers: {', '.join(unavailable_providers)}. Please check docs page if you wish to enable them."
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Filter model types based on provider availability in MODEL_CLASS_MAP
|
available_model_types = esperanto_available_providers.keys()
|
||||||
available_model_types = []
|
|
||||||
for model_type in model_types:
|
|
||||||
if model_type in MODEL_CLASS_MAP and provider in MODEL_CLASS_MAP[model_type]:
|
|
||||||
available_model_types.append(model_type)
|
|
||||||
|
|
||||||
if not available_model_types:
|
|
||||||
st.error(f"No compatible model types available for provider: {provider}")
|
|
||||||
else:
|
|
||||||
model_type = st.selectbox(
|
model_type = st.selectbox(
|
||||||
"Model Type",
|
"Model Type",
|
||||||
available_model_types,
|
available_model_types,
|
||||||
help="Use language for text generation models, text_to_speech for TTS models for generating podcasts, etc.",
|
help="Use language for text generation models, text_to_speech for TTS models for generating podcasts, etc.",
|
||||||
)
|
)
|
||||||
|
provider = st.selectbox("Provider", esperanto_available_providers[model_type])
|
||||||
|
|
||||||
if model_type == "text_to_speech" and provider == "gemini":
|
if model_type == "text_to_speech" and provider == "gemini":
|
||||||
model_name = "gemini-default"
|
model_name = "gemini-default"
|
||||||
st.markdown("Gemini models are pre-configured. Using the default model.")
|
st.markdown("Gemini models are pre-configured. Using the default model.")
|
||||||
|
|
@ -140,6 +146,8 @@ with model_tab:
|
||||||
new_model = Model(**recommendation)
|
new_model = Model(**recommendation)
|
||||||
new_model.save()
|
new_model.save()
|
||||||
st.rerun()
|
st.rerun()
|
||||||
|
st.divider()
|
||||||
|
|
||||||
st.subheader("Configured Models")
|
st.subheader("Configured Models")
|
||||||
model_types_available = {
|
model_types_available = {
|
||||||
# "vision": False,
|
# "vision": False,
|
||||||
|
|
@ -160,7 +168,10 @@ with model_tab:
|
||||||
if not available:
|
if not available:
|
||||||
st.warning(f"No models available for {model_type}")
|
st.warning(f"No models available for {model_type}")
|
||||||
|
|
||||||
with model_defaults_tab:
|
|
||||||
|
st.divider()
|
||||||
|
|
||||||
|
st.subheader("Select Default Models")
|
||||||
text_generation_models = [model for model in all_models if model.type == "language"]
|
text_generation_models = [model for model in all_models if model.type == "language"]
|
||||||
|
|
||||||
text_to_speech_models = [
|
text_to_speech_models = [
|
||||||
|
|
@ -254,7 +265,7 @@ with model_defaults_tab:
|
||||||
st.divider()
|
st.divider()
|
||||||
# Handle embedding model selection
|
# Handle embedding model selection
|
||||||
selected_model = model_selector(
|
selected_model = model_selector(
|
||||||
"Default Speech to Text Model",
|
"Default Embedding Model",
|
||||||
"default_embedding_model",
|
"default_embedding_model",
|
||||||
selected_id=default_models.default_embedding_model,
|
selected_id=default_models.default_embedding_model,
|
||||||
help="This is the default model for embeddings (semantic search, etc)",
|
help="This is the default model for embeddings (semantic search, etc)",
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue