Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions astrbot/core/config/default.py
Original file line number Diff line number Diff line change
Expand Up @@ -1904,9 +1904,9 @@
"enable": True,
"embedding_api_key": "",
"embedding_api_base": "https://integrate.api.nvidia.com/v1",
"embedding_model": "nvidia/llama-nemotron-embed-1b-v2",
"embedding_model": "nvidia/nemotron-3-embed-1b",
"input_type": "passage",
"embedding_dimensions": 1024,
"embedding_dimensions": 2048,
"timeout": 20,
"proxy": "",
},
Expand Down Expand Up @@ -1982,7 +1982,7 @@
"enable": True,
"nvidia_rerank_api_key": "",
"nvidia_rerank_api_base": "https://ai.api.nvidia.com/v1/retrieval",
"nvidia_rerank_model": "nv-rerank-qa-mistral-4b:1",
"nvidia_rerank_model": "nvidia/llama-nemotron-rerank-vl-1b-v2",
"nvidia_rerank_model_endpoint": "/reranking",
"timeout": 20,
"nvidia_rerank_truncate": "",
Expand Down
2 changes: 1 addition & 1 deletion astrbot/core/provider/sources/nvidia_embedding_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ def __init__(self, provider_config: dict, provider_settings: dict) -> None:
)
self.timeout = int(provider_config.get("timeout", 20))
self.model = provider_config.get(
"embedding_model", "nvidia/llama-nemotron-embed-1b-v2"
"embedding_model", "nvidia/nemotron-3-embed-1b"
)
self.input_type = provider_config.get("input_type", "passage")

Expand Down
2 changes: 1 addition & 1 deletion astrbot/core/provider/sources/nvidia_rerank_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ def __init__(self, provider_config: dict, provider_settings: dict) -> None:
).rstrip("/")
self.timeout = provider_config.get("timeout", 20)
self.model = provider_config.get(
"nvidia_rerank_model", "nv-rerank-qa-mistral-4b:1"
"nvidia_rerank_model", "nvidia/llama-nemotron-rerank-vl-1b-v2"
)
self.model_endpoint = provider_config.get(
"nvidia_rerank_model_endpoint", "/reranking"
Expand Down
61 changes: 61 additions & 0 deletions tests/test_nvidia_embedding_source.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
from astrbot.core.config.default import CONFIG_METADATA_2
from astrbot.core.provider.sources.nvidia_embedding_source import (
NvidiaEmbeddingProvider,
)

NEW_MODEL = "nvidia/nemotron-3-embed-1b"
OLD_MODEL = "nvidia/llama-nemotron-embed-1b-v2"


def test_nvidia_embedding_config_template_uses_new_model_and_dimension():
templates = CONFIG_METADATA_2["provider_group"]["metadata"]["provider"][
"config_template"
]

assert templates["NVIDIA Embedding"]["embedding_model"] == NEW_MODEL
assert templates["NVIDIA Embedding"]["embedding_dimensions"] == 2048


def test_nvidia_embedding_provider_uses_new_fallback_model():
provider = NvidiaEmbeddingProvider({}, {})

assert provider.model == NEW_MODEL
assert provider.get_model() == NEW_MODEL


def test_nvidia_embedding_provider_preserves_explicit_old_model():
provider = NvidiaEmbeddingProvider(
{
"embedding_model": OLD_MODEL,
"embedding_dimensions": 1024,
},
{},
)

assert provider.model == OLD_MODEL
Comment thread
sourcery-ai[bot] marked this conversation as resolved.
assert provider.get_dim() == 1024


def test_nvidia_embedding_new_model_uses_existing_api_contract():
provider = NvidiaEmbeddingProvider(
{
"embedding_model": NEW_MODEL,
"input_type": "passage",
},
{},
)

assert provider._build_payload(["first", "second"]) == {
"input": ["first", "second"],
"model": NEW_MODEL,
"input_type": "passage",
"encoding_format": "float",
}
assert provider._parse_response(
{
"data": [
{"index": 0, "embedding": [0.1, 0.2]},
{"index": 1, "embedding": [0.3, 0.4]},
]
}
) == [[0.1, 0.2], [0.3, 0.4]]
69 changes: 69 additions & 0 deletions tests/test_nvidia_rerank_source.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,69 @@
from astrbot.core.config.default import CONFIG_METADATA_2
from astrbot.core.provider.sources.nvidia_rerank_source import NvidiaRerankProvider

NEW_MODEL = "nvidia/llama-nemotron-rerank-vl-1b-v2"
OLD_MODEL = "nv-rerank-qa-mistral-4b:1"


def test_nvidia_rerank_config_template_uses_new_model():
templates = CONFIG_METADATA_2["provider_group"]["metadata"]["provider"][
"config_template"
]

assert templates["NVIDIA Rerank"]["nvidia_rerank_model"] == NEW_MODEL


def test_nvidia_rerank_provider_uses_new_fallback_model():
provider = NvidiaRerankProvider({}, {})

assert provider.model == NEW_MODEL
assert provider.get_model() == NEW_MODEL


def test_nvidia_rerank_provider_preserves_explicit_old_model():
provider = NvidiaRerankProvider({"nvidia_rerank_model": OLD_MODEL}, {})

assert provider.model == OLD_MODEL
assert provider._get_endpoint() == (
"https://ai.api.nvidia.com/v1/retrieval/nvidia/reranking"
)


def test_nvidia_rerank_new_model_uses_existing_api_contract():
provider = NvidiaRerankProvider(
{
"nvidia_rerank_model": NEW_MODEL,
"nvidia_rerank_truncate": "END",
},
{},
)

assert provider._get_endpoint() == (
"https://ai.api.nvidia.com/v1/retrieval/nvidia/"
"llama-nemotron-rerank-vl-1b-v2/reranking"
)
assert provider._build_payload("query", ["first", "second"]) == {
"model": NEW_MODEL,
"query": {"text": "query"},
"passages": [{"text": "first"}, {"text": "second"}],
"truncate": "END",
}


def test_nvidia_rerank_parses_official_rankings_response():
provider = NvidiaRerankProvider({}, {})

results = provider._parse_results(
{
"rankings": [
{"index": 1, "logit": -0.25},
{"index": 0, "logit": 0.75},
]
},
top_n=None,
)

assert [(result.index, result.relevance_score) for result in results] == [
(0, 0.75),
(1, -0.25),
]
Loading