文件導航

重排序

為 40 條 CLERC 法律查詢各構建 30 個段落的 BM25 候選列表,然後為每個查詢-候選對提一個 TypeSafe 問題,把 top-1 準確率從 5% 提升到 18%,top-10 從 38% 提升到 62%。

你有成千上萬份文件,需要找到能回答某個具體問題的那一份。那該怎麼找?

首先,用一種快速的方法,比如關鍵詞匹配,把成千上萬的候選縮小成一份看起來靠譜的候選列表。我們稱之為快速搜尋。

快速搜尋擅長這件事,但它無法告訴你候選列表裡哪一個才是對的。這正是重排序的用武之地。它把候選列表上的每個候選直接對照查詢打分,把最好的排到最前面。

下面兩步都會在 CLERC 資料集的 3,565 個法院意見段落上執行:BM25 為 40 條查詢各構建一份 30 個候選的快速搜尋列表,然後 TypeSafe 對每份列表重排序。經過重排序,正確的段落在 18% 的查詢上排到第一,而只用快速搜尋時是 5%。

沿途你會學到:

  • 快速搜尋做什麼,以及它為什麼不是全部答案
  • 重排序是什麼,以及它如何接在快速搜尋之後
  • TypeSafe 如何把某個候選用查詢來打分,以及這能把結果提升多少

自己試試

在 TypeSafe Playground 中開啟一條查詢、一個候選和重排序問題

如何在成千上萬份文件裡找到那一份?

你有一堆文件,還有一個查詢——一段描述你要找什麼的文字。這堆文件裡的某處,就有能回答它的那一份。

逐個把每份文件和查詢比對是可行的,但每份文件要一次比對:幾百萬份文件就意味著每條查詢幾百萬次比對。你可以用一個兩步走的方法來提升效能:

  1. 用一種快到能跑遍整堆文件的方法,把這堆縮小成一份可能候選的短列表。
  2. 對這份短列表施加一個更精確的步驟,找出確切正確的答案。
動畫圖示:一堆文件收窄成一份快速搜尋候選列表,隨後重排序重新排列這份列表,讓正確答案升到最前面

這個 cookbook 用一批法院意見資料集測試了這套設定,見下面的一個重排序示例。

什麼是快速搜尋?

快速搜尋是任何能把查詢和大型語料庫中的每份文件比對、並快速返回一份排好序的短列表的方法。常見方法包括關鍵詞搜尋(如 BM25)和按含義比較段落的稠密嵌入。系統常常把兩者結合起來。

這裡第一步只用 BM25,別的都不用。BM25 按共有的詞給段落排序。讓這一步保持簡單,是為了把注意力留在重排序上,這正是本 cookbook 的重點。快速搜尋方法的選擇是次要問題:重排序只能看到進入短列表的那些段落。

什麼是重排序?

重排序拿到快速搜尋已經產出的短列表,把它排成更好的順序。它不是一次性把查詢和整個語料庫比對,而是把查詢逐個對照短列表上的每個候選,再按這個分數給短列表排序。

圖示:左邊是一份排好序的短列表,中間一個標著“re-rank”的箭頭,右邊是重排後的版本,真正的答案從中間移到最前面

分數可以來自語言模型。把查詢和一個候選一起給它,問這個候選對查詢的回答有多好。這樣,即使候選的措辭和查詢不同,重排序也能在短列表上找到最佳匹配。

用 TypeSafe 重排序

重排序器需要為每個查詢-候選對給出可比較的分數。通用語言模型可以產生這些分數,也可以直接給整份短列表排序。但對於獨立的成對打分,你需要定義一個打分尺度,並提示模型對每個候選應用同一標準。重複呼叫對同一個對仍可能給出不同的分數,而通用生成會給一個只需要一個數字的任務增加時間和成本。

TypeSafe 返回什麼

用 TypeSafe,打分請求可以保持是一個是/否問題:

Could this candidate passage be from the cited precedent?

光是一個是或否,不足以給 30 個候選排序。Noul 會返回一個 0 到 1 之間的數字,叫作 noul。noul 是 TypeSafe 對“答案有多可能是是”的估計。

問題的 criteria 定義了什麼算真、什麼算假。TypeSafe 把它們應用到每個查詢-候選對上,直接返回 noul。這個 noul 就是應用用來排序的分數。不必再為通用模型發明一套打分尺度,而 TypeSafe 正是為更快、更省、更一致地做這種重複打分而構建的。

簡化成虛擬碼,一次 TypeSafe 打分呼叫大致是這樣:

question = Noul(
    instructions="Is this candidate the cited case?",
    criteria=NoulCriteria(
        true="The candidate states the specific rule the query cites.",
        false="The candidate is only on a similar topic.",
    ),
)
response = client.system_one(state={...}, questions={"is_cited_source": question})
response.answers["is_cited_source"].noul  # -> 0.87

TypeSafe 把查詢和一個候選放在一起對照那個問題來讀,返回一個 noul。

你可以用它重排序一份短列表:對列表上的每個候選跑同一個問題,然後按每次呼叫返回的 noul 給短列表排序,最高的在前。

nouls = {candidate: ask_typesafe(query, candidate) for candidate in shortlist}
reranked = sorted(shortlist, key=lambda c: nouls[c], reverse=True)  # highest noul first

下圖展示了每個候選一個請求如何產生用來重排短列表的分數。

flowchart LR
    q["query excerpt<br/><i>one opinion passage,<br/>citation removed</i>"]
    sl["shortlist from fast search<br/><i>30 candidate passages</i>"]
    quest["<b>one Noul</b><br/>could this candidate be<br/>from the cited precedent?<br/><i>criteria fix true and false</i>"]

    %% direction LR inside an LR chart keeps each state beside its noul, two columns,
    %% so the fan-out is four rows tall instead of eight
    subgraph fan["one request per candidate · no request sees another"]
        direction LR
        d1["state<br/>{query, candidate 1}"] --> n1["noul<br/>0.87"]
        d2["state<br/>{query, candidate 2}"] --> n2["noul<br/>0.41"]
        dx["⋮"] --> nx["⋮"]
        d30["state<br/>{query, candidate 30}"] --> n30["noul<br/>0.12"]
    end

    sort["sort by noul,<br/>highest first"]
    out["re-ranked shortlist<br/><i>same 30, better order</i>"]

    q --> fan
    sl --> fan
    quest --> fan
    fan --> sort --> out

    %% the elision is not a node - drop its box so it reads as "and so on"
    classDef elide fill:none,stroke:none
    class dx,nx elide
    linkStyle 2 stroke:none

一個重排序示例

快速搜尋和重排序現在跑在 CLERC 上,這是一個法律檢索資料集。本例使用 3,565 個法院意見段落和 40 條查詢。

環境準備

第一步安裝本演示依賴的包。

  • bm25s 和 datasets 構建快速搜尋的短列表。
  • typesafe-sdk 和 cooksafe 負責重排序和 API 快取。
  • matplotlib 繪製結果圖表。
pip install bm25s datasets matplotlib 'cooksafe>=0.2.0,<0.3.0'

下一個程式碼塊設定 TypeSafe 客戶端,以及本演示其餘部分用到的常量,比如呼叫哪個 TypeSafe 模型、快速搜尋交給重排序器的短列表有多大。呼叫 TypeSafe 需要一個 TYPESAFE_API_KEY。

import hashlib
import json
import os
import random
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path

from cooksafe import JsonCache
from IPython.display import display
from typesafe_sdk import Noul, NoulCriteria, TypeSafeClient

TYPESAFE_MODEL = "jev-1.12"
PRICE = (
    0.042,
    0.00,
)  # $ per 1M tokens (input, output); TypeSafe jev-1.12 as of 2026-08
N_ROWS = 170  # CLERC rows pooled into the shared corpus
N_QUERIES = 40  # rows we evaluate
TOP_K = 30  # candidates the shortlist hands to the re-ranker, per query

client = TypeSafeClient(
    api_key=os.environ.get(
        "TYPESAFE_API_KEY", "cache-only"
    ),  # keyless kernels replay the cache
    base_url=os.environ.get("TYPESAFE_ENDPOINT"),
    timeout=120.0,
)
json_cache = JsonCache(Path("json_cache.json"))

用快速搜尋給段落排序

這裡用的資料集是美國法院意見的語料庫,共彙總了 170 行。每一行拆開來看是這樣:

  • 查詢:一段去掉了引用的意見摘錄。
  • 金標準:被去掉的那條引用所指向的段落,也就是查詢的唯一正確答案。
  • 候選:語料庫中所有其它段落,每一個都是查詢可能被誤匹配的物件。

170 行裡有 40 行被選出來作為查詢評估。其餘 130 行只作為候選出現。

下一個單元格用上面描述的方法構建短列表:

  1. 載入語料庫。
  2. 用 BM25 把語料庫對每條查詢排序。

這裡還沒有 TypeSafe,這僅僅是快速搜尋這一步。

CLERC_FILE = (
    "https://huggingface.co/datasets/jhu-clsp/CLERC/resolve/main/"
    "teva_train_dir/train_data.jsonl.gz"
)

def cid(text: str) -> str:
    """Corpus id: a content hash, so passages shared across queries dedupe."""
    return hashlib.sha1(text.encode("utf-8")).hexdigest()[:16]

@json_cache
def build_slice(n_rows: int, n_queries: int, seed: int) -> dict:
    """Stream CLERC rows, pool ``n_rows`` of them into a corpus, pick ``n_queries`` to evaluate."""
    from datasets import load_dataset  # heavy import, keep local

    stream = load_dataset("json", data_files=CLERC_FILE, streaming=True, split="train")
    rows = []
    for row in stream:
        if (
            row.get("positive_passages")
            and len(row.get("negative_passages") or []) == 20
        ):
            rows.append(row)
        if len(rows) >= 1000:
            break

    rng = random.Random(seed)
    picked = rng.sample(rows, n_rows)
    corpus, pool = {}, []
    for row in picked:
        gold = row["positive_passages"][0]["text"]
        corpus[cid(gold)] = gold
        for neg in row["negative_passages"]:
            corpus[cid(neg["text"])] = neg["text"]
        pool.append(
            {"qid": str(row["query_id"]), "query": row["query"], "gold": cid(gold)}
        )
    # hold out the first 20 pooled rows; evaluate on the rest
    queries = rng.sample(pool[20:], n_queries)
    # sort the corpus by id so every run — live or cache replay — iterates it identically
    return {"queries": queries, "corpus": dict(sorted(corpus.items()))}

def bm25_rankings(corpus: dict[str, str], queries: dict[str, str], k: int = 100):
    """Rank every passage in the corpus by word overlap with each query."""
    import bm25s

    cids = list(corpus)
    retriever = bm25s.BM25()
    retriever.index(bm25s.tokenize([corpus[c] for c in cids], stopwords="en"))
    qids = list(queries)
    idxs, _ = retriever.retrieve(
        bm25s.tokenize([queries[q] for q in qids], stopwords="en"), k=min(k, len(cids))
    )
    return {q: [cids[i] for i in idxs[row]] for row, q in enumerate(qids)}

def gold_rank(ranked: list[str], gold: str) -> int | None:
    """1-based rank of the gold id, or None if it isn't in the list."""
    return ranked.index(gold) + 1 if gold in ranked else None

SURFACE, INK, INK2, MUTED = "#f8f8f2", "#34342f", "#34342f", "#7c7c77"
GRID, AXIS, BLUE, GREEN = "#d8d8cf", "#d8d8cf", "#5d76a2", "#6f9b52"

def bar_chart(labels: list[str], shares: list[float], title: str) -> None:
    """A small single-series bar chart of shares (0-1, shown as percentages)."""
    import matplotlib.pyplot as plt

    fig, ax = plt.subplots(figsize=(5, 3.2), facecolor=SURFACE)
    ax.set_facecolor(SURFACE)
    for side in ("top", "right"):
        ax.spines[side].set_visible(False)
    for side in ("left", "bottom"):
        ax.spines[side].set_color(AXIS)
    ax.tick_params(colors=MUTED, labelcolor=INK2, labelsize=9)
    ax.set_axisbelow(True)
    ax.grid(axis="y", color=GRID, linewidth=0.8)

    bars = ax.bar(labels, shares, width=0.55, color=[BLUE, GREEN][: len(labels)])
    ax.bar_label(
        bars,
        labels=[f"{s * 100:.0f}%" for s in shares],
        padding=4,
        color=INK,
        fontsize=11,
    )
    ax.set_ylim(0, 1.1)
    ax.set_yticks([0, 0.25, 0.5, 0.75, 1.0])
    ax.set_yticklabels(["0%", "25%", "50%", "75%", "100%"])
    ax.set_ylabel(f"share of {len(queries)} queries", color=INK2, fontsize=9)
    ax.set_title(title, loc="left", color=INK, fontsize=11)
    plt.tight_layout()
    display(fig)
    plt.close(fig)

ds = build_slice(N_ROWS, N_QUERIES, seed=0)
corpus: dict[str, str] = ds["corpus"]
queries = {q["qid"]: q["query"] for q in ds["queries"]}
golds = {q["qid"]: q["gold"] for q in ds["queries"]}

candidates = {q: ranked[:TOP_K] for q, ranked in bm25_rankings(corpus, queries).items()}

in_top_k = sum(golds[q] in candidates[q] for q in queries)
at_rank_1 = sum(candidates[q][0] == golds[q] for q in queries)

bar_chart(
    [f"In top {TOP_K}", "At rank 1"],
    [in_top_k / len(queries), at_rank_1 / len(queries)],
    f"Where the correct passage lands, {len(queries)} queries against {len(corpus):,} candidates",
)
輸出

快速搜尋不太可能把正確的段落排在第一

這張圖顯示在 3,565 個候選中,快速搜尋把正確段落放在哪裡。

快速搜尋能可靠地把語料庫縮小到一份包含正確答案的短列表。對於 40 條查詢,它 100% 都包含正確答案。但那個段落很少是短列表上排第一的,只有 5% 的時候是。

下面的重排序只重新排列短列表上已有的前 30 個候選。它無法加入快速搜尋沒有選中的段落。在這裡,短列表對全部 40 條查詢都包含正確的段落,所以重排序可以專注於把每一個放到更好的位置。

用 TypeSafe 重排序

重排序把短列表上每個候選對照它的查詢打分,再按這個分數排序。TypeSafe 對每一對問的問題是:這個候選有沒有可能是查詢那條被去掉的引用所指向的段落。

下一個單元格做以下事情:

  1. 定義那個問題。
  2. 對每份短列表上的每個候選各問一次,40 條查詢乘以 30 個候選,共 1,200 次呼叫,併發執行而不是一個接一個。
  3. 按 TypeSafe 返回的分數給每份短列表排序,得到重排序後的結果。
is_cited_source = Noul(
    instructions=(
        "The query excerpt comes from a US federal court opinion and was written "
        "immediately around a citation to a precedent; the citation itself has been "
        "removed. Could the candidate passage be from that cited precedent — does it "
        "establish the specific legal proposition the query excerpt invokes at its "
        "citation point?"
    ),
    criteria=NoulCriteria(
        true=(
            "The candidate passage states or establishes the specific rule, standard, "
            "holding, or fact pattern that the query excerpt attributes to its removed "
            "citation."
        ),
        false=(
            "The candidate passage is merely on a similar topic or doctrine; it does not "
            "supply the specific proposition the query excerpt relies on."
        ),
    ),
)

@json_cache
def score_candidate(model: str, query: str, candidate: str, question_json: str) -> dict:
    """One TypeSafe call about one (query, candidate) pair: a noul, plus token usage."""
    # the SDK takes a question as its JSON dict, so the cached string decodes straight in
    question = json.loads(question_json)
    response = client.system_one(
        state={"query_excerpt": query, "candidate_passage": candidate},
        questions={"is_cited_source": question},
        model=model,
    )
    return {
        "noul": response.answers["is_cited_source"].noul,
        "input_tokens": response.usage.input_tokens or 0,
        "output_tokens": response.usage.output_tokens or 0,
    }

# Each of the 40 queries has 30 candidates, so re-ranking every shortlist means 1,200 independent
# calls — cheap enough to fire all at once with a thread pool instead of one after another.
pair_list = [(q, c) for q in queries for c in candidates[q]]
question_json = is_cited_source.model_dump_json(exclude_none=True)
with ThreadPoolExecutor(max_workers=12) as pool:
    results = pool.map(
        lambda p: score_candidate(
            TYPESAFE_MODEL, queries[p[0]], corpus[p[1]], question_json
        ),
        pair_list,
    )
pair_scores = {q: {} for q in queries}
for (q, c), result in zip(pair_list, results):
    pair_scores[q][c] = result

reranked = {
    q: sorted(candidates[q], key=lambda c: -pair_scores[q][c]["noul"]) for q in queries
}

def chart_before_after(
    runs: dict[str, dict[str, list[str]]], thresholds: list[int]
) -> None:
    """Grouped bar chart: how often the correct passage lands in the top N, for each run."""
    import numpy as np
    import matplotlib.pyplot as plt

    labels = list(runs)
    colors = [BLUE, GREEN]

    def share_in_top(rankings, k):
        return sum(
            gold_rank(rankings[q], golds[q]) in range(1, k + 1) for q in queries
        ) / len(queries)

    fig, ax = plt.subplots(figsize=(6.5, 3.6), facecolor=SURFACE)
    ax.set_facecolor(SURFACE)
    for side in ("top", "right"):
        ax.spines[side].set_visible(False)
    for side in ("left", "bottom"):
        ax.spines[side].set_color(AXIS)
    ax.tick_params(colors=MUTED, labelcolor=INK2, labelsize=9)
    ax.set_axisbelow(True)
    ax.grid(axis="y", color=GRID, linewidth=0.8)

    x = np.arange(len(thresholds))
    width = 0.35
    for i, (label, rankings) in enumerate(runs.items()):
        shares = [share_in_top(rankings, k) for k in thresholds]
        offset = (i - (len(labels) - 1) / 2) * width
        bars = ax.bar(x + offset, shares, width * 0.92, color=colors[i], label=label)
        ax.bar_label(
            bars,
            labels=[f"{s * 100:.0f}%" for s in shares],
            padding=3,
            color=INK2,
            fontsize=8.5,
        )

    ax.set_xticks(x, [f"top {k}" for k in thresholds])
    ax.set_ylim(0, 1)
    ax.set_yticks([0, 0.25, 0.5, 0.75, 1.0])
    ax.set_yticklabels(["0%", "25%", "50%", "75%", "100%"])
    ax.set_ylabel(f"share of {len(queries)} queries", color=INK2, fontsize=9)
    ax.set_title(
        "How often the correct passage lands near the top",
        loc="left",
        color=INK,
        fontsize=11,
    )
    ax.legend(frameon=False, labelcolor=INK2, fontsize=9, loc="upper left")
    plt.tight_layout()
    display(fig)
    plt.close(fig)

chart_before_after(
    {"Fast search": candidates, "+ TypeSafe re-rank": reranked}, [1, 5, 10]
)

calls = [pair_scores[q][c] for q in queries for c in pair_scores[q]]
input_tokens = sum(call["input_tokens"] for call in calls)
output_tokens = sum(call["output_tokens"] for call in calls)
cost = input_tokens / 1_000_000 * PRICE[0] + output_tokens / 1_000_000 * PRICE[1]
print(
    f"{len(calls)} TypeSafe calls used {input_tokens:,} input and "
    f"{output_tokens:,} output tokens, costing ${cost:.4f}."
)
1200 TypeSafe calls used 1,536,002 input and 25,200 output tokens, costing $0.0645.
輸出

重排序把正確答案移向最前面

這張圖在三個閾值上比較快速搜尋和快速搜尋加重排序。在每一個閾值上,重排序都把正確的段落移得更靠近最前面:

  • Top 1 — 5% → 18%
  • Top 5 — 15% → 35%
  • Top 10 — 38% → 62%

報告的 token 數和成本涵蓋用於重排序這 40 份短列表的全部 1,200 次 TypeSafe 呼叫。

每一行 CLERC 資料包含一個正確段落和 20 個負例段落。本演示把 170 行的段落彙總進一個共享語料庫。對 40 條評估查詢中的每一條,BM25 從整個語料庫中選出 30 個候選,而不只是該行附帶的 20 個負例。隨後 TypeSafe 把查詢對照每個選中的候選來讀,並重排序那 30 個段落。

為了清晰,本演示對每一對只問了一個問題。真實應用會在一次呼叫中對同一對問好幾個問題。做法見並行問題 cookbook和投機式扇出模式。


接下來

同樣的積木也出現在 TypeSafe 文件的其它地方:

  • Noul,講 TypeSafe 如何把一個是/否問題變成一個分數。
  • 投機式扇出,講如何在一次呼叫中對一份文件問好幾個問題。
  • 逐行搜尋,講另一種按含義而不是關鍵詞檢索語料庫的方式。