Text Ranking
sentence-transformers
Safetensors
Transformers
multilingual
t5gemma2
text2text-generation
reranker
encoder-decoder
FBNL
Retrieval
RAG
Instructions to use KaLM-Embedding/KaLM-Reranker-V1-Large with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- sentence-transformers
How to use KaLM-Embedding/KaLM-Reranker-V1-Large with sentence-transformers:
from sentence_transformers import CrossEncoder model = CrossEncoder("KaLM-Embedding/KaLM-Reranker-V1-Large") query = "Which planet is known as the Red Planet?" passages = [ "Venus is often called Earth's twin because of its similar size and proximity.", "Mars, known for its reddish appearance, is often referred to as the Red Planet.", "Jupiter, the largest planet in our solar system, has a prominent red spot.", "Saturn, famous for its rings, is sometimes mistaken for the Red Planet." ] scores = model.predict([(query, passage) for passage in passages]) print(scores) - Transformers
How to use KaLM-Embedding/KaLM-Reranker-V1-Large with Transformers:
# Load model directly from transformers import AutoProcessor, AutoModelForMultimodalLM processor = AutoProcessor.from_pretrained("KaLM-Embedding/KaLM-Reranker-V1-Large") model = AutoModelForMultimodalLM.from_pretrained("KaLM-Embedding/KaLM-Reranker-V1-Large", device_map="auto") - Notebooks
- Google Colab
- Kaggle
File size: 3,914 Bytes
6f7a484 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 | from __future__ import annotations
import argparse
import json
import sys
import urllib.error
import urllib.request
from pathlib import Path
from typing import Any
from .constants import SAMPLE_DOCUMENTS, SAMPLE_QUERY
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description="Call the KaLM vLLM FastAPI service.",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument("--base-url", default="http://127.0.0.1:8000")
parser.add_argument("--health", action="store_true")
parser.add_argument("--endpoint", choices=("rerank", "score"), default="rerank")
parser.add_argument("--json-file", type=Path)
parser.add_argument("--return-margin", action="store_true")
parser.add_argument("--top-k", type=int)
parser.add_argument("--timeout", type=float, default=600.0)
return parser
def request_json(
method: str,
url: str,
*,
payload: dict[str, Any] | None = None,
timeout: float = 600.0,
) -> dict[str, Any]:
data = None
headers = {"Accept": "application/json"}
if payload is not None:
data = json.dumps(payload, ensure_ascii=False).encode("utf-8")
headers["Content-Type"] = "application/json"
request = urllib.request.Request(url, data=data, headers=headers, method=method)
opener = urllib.request.build_opener(urllib.request.ProxyHandler({}))
try:
with opener.open(request, timeout=timeout) as response:
raw = response.read().decode("utf-8")
except urllib.error.HTTPError as error:
detail = error.read().decode("utf-8", errors="replace")
raise RuntimeError(f"HTTP {error.code} from {url}: {detail}") from error
return json.loads(raw)
def _demo_payload(
endpoint: str,
return_margin: bool,
top_k: int | None,
) -> dict[str, Any]:
if endpoint == "rerank":
return {
"query": SAMPLE_QUERY,
"documents": list(SAMPLE_DOCUMENTS),
"top_k": top_k,
"return_margin": return_margin,
}
return {
"pairs": [
{
"id": "positive",
"query": SAMPLE_QUERY,
"document": SAMPLE_DOCUMENTS[0],
},
{
"id": "negative",
"query": SAMPLE_QUERY,
"document": SAMPLE_DOCUMENTS[1],
},
],
"return_margin": return_margin,
}
def _load_payload(path: Path) -> dict[str, Any]:
with path.open("r", encoding="utf-8") as handle:
payload = json.load(handle)
if not isinstance(payload, dict):
raise ValueError(f"{path} must contain one JSON object.")
return payload
def _request_payload(args: argparse.Namespace) -> dict[str, Any]:
payload = (
_load_payload(args.json_file)
if args.json_file is not None
else _demo_payload(args.endpoint, args.return_margin, args.top_k)
)
if args.return_margin:
payload["return_margin"] = True
if args.top_k is not None:
if args.endpoint != "rerank":
raise ValueError("--top-k is only valid with --endpoint rerank.")
if args.top_k < 0:
raise ValueError("--top-k must be non-negative.")
payload["top_k"] = args.top_k
return payload
def main() -> int:
args = build_parser().parse_args()
base_url = args.base_url.rstrip("/")
if args.health:
response = request_json("GET", f"{base_url}/health", timeout=args.timeout)
else:
payload = _request_payload(args)
response = request_json(
"POST",
f"{base_url}/{args.endpoint}",
payload=payload,
timeout=args.timeout,
)
json.dump(response, sys.stdout, ensure_ascii=False, indent=2)
sys.stdout.write("\n")
return 0
if __name__ == "__main__":
raise SystemExit(main())
|