Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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.
Outdated
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