TS TypeSafe 文档中文版 原文 ↗

重排序

为 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. 对这份短列表应用一个更精确的步骤,找出确切正确的答案。

动画示意图:一堆文档收窄成一份快速搜索候选列表,接着重排序
重新排列这份候选列表,使正确答案上升到
顶部

本实践手册用一份法院判决书数据集来检验这套设置, 见下文 在一个真实例子上做重排序。

快速搜索是任何能把查询与大型语料库中的每份文档比对、 并快速返回一份带排序的候选列表的方法。常见方法包括关键词搜索(如 BM25) 和按语义比较段落的稠密嵌入。系统常常把两种 方法结合起来。

这里的第一步只有 BM25,没有别的。BM25 按共有词对段落排序。 让这一步保持简单,才能把注意力留给重排序,而重排序正是本实践手册的 重点。选择哪种快速搜索方法是个次要问题:重排序看到的永远只是 那些进入候选列表的段落。

什么是重排序?

重排序接过快速搜索已经产出的候选列表,把它排成更好的顺序。 它不是一次性把查询与整个语料库比对,而是把查询 与候选列表上的每个候选项逐一比对,再按该 得分对候选列表排序。

示意图:左侧是一份带排序的候选列表,一个标有 "重排序" 的箭头,
右侧是重新排序后的版本,真实答案从中部移到
顶部

得分可以来自语言模型。把查询和一个候选项一起交给它, 问这个候选项在多大程度上回答了查询。这样,重排序就能在候选列表上找到 最佳匹配,即使它的措辞与查询不同。

用 TypeSafe 做重排序

重排序器需要为每个查询-候选对给出可比较的得分。通用 语言模型可以产出这些得分,也可以直接对整个候选列表排序。但如果 要独立地为每个对打分,你就需要定义一个评分尺度,并提示模型 对每个候选项都应用同一标准。重复调用仍可能对同一个对给出不同的 得分,而通用生成又会给一个只需要一个数字的任务 增加时间和成本。

TypeSafe 返回什么

有了 TypeSafe,评分请求可以就是一个是/否问题:

plaintext
Could this candidate passage be from the cited precedent?
 

一个简单的“是”或“否”不足以给 30 个候选项排序。Noul 则 返回一个 0 到 1 之间的数字,称为 noul。noul 是 TypeSafe 对 答案有多大可能是“是”的估计。

问题的 criteria 定义了什么算真、什么算假。TypeSafe 把它们应用到 每个查询-候选对上,并直接返回 noul。这个 noul 就是应用 用来排序的得分。不必再为通用模型发明一套评分尺度,而 TypeSafe 正是为了更快、更便宜、更一致地做这种重复评分而构建的。

用简化的伪代码表示,一次 TypeSafe 评分调用是这样的:

python
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 对候选列表排序, 最高的在前。

python
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["查询摘录<br/><i>一段判决书段落,<br/>已移除引用</i>"] sl["来自快速搜索的候选列表<br/><i>30 个候选段落</i>"] quest["<b>一个 Noul</b><br/>这个候选项会不会<br/>来自被引用的先例?<br/><i>criteria 确定真与假</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["每个候选项一个请求 · 没有请求能看到其他请求"] direction LR d1["状态<br/>{query, candidate 1}"] --> n1["noul<br/>0.87"] d2["状态<br/>{query, candidate 2}"] --> n2["noul<br/>0.41"] dx["⋮"] --> nx["⋮"] d30["状态<br/>{query, candidate 30}"] --> n30["noul<br/>0.12"] end sort["按 noul 排序,<br/>最高在前"] out["重排序后的候选列表<br/><i>同样 30 个,更好的顺序</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 用于绘制结果图表。
bash
pip install bm25s datasets matplotlib "typesafe-sdk>=0.5.7" cooksafe --extra-index-url https://pypi.typesafe.ai/
 

下一个代码块设置 TypeSafe 客户端,以及本次演练其余部分 用到的常量,例如调用哪个 TypeSafe 模型、快速搜索交给 重排序器的候选列表有多大。调用 TypeSafe 需要一个 TYPESAFE_API_KEY。

python
import hashlib
import json
import os
import random
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
 
import msgspec
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 行汇总而成。每一 行的构成如下:

  • 查询(Query):一段移除了引用的判决书摘录。
  • Gold:被移除的引用所指向的段落,也就是该查询唯一正确的 答案。
  • 候选项(Candidates):语料库中其余所有段落,每一个都可能被查询 错误地匹配上。

在这 170 行中,挑出 40 行作为查询来评估。其余 130 行只作为 候选项出现。

下一个单元格用上文描述的技术构建候选列表:

  1. 加载语料库。
  2. 用 BM25 针对每个查询对它排序。

这里还没有用到 TypeSafe,这只是快速搜索这一步。

python
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 返回的得分对每份候选列表排序,得到重排序后的结果。
python
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 = msgspec.json.encode(is_cited_source).decode()
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}."
)
 
text
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 个段落重排序。

为了清晰起见,本演练对每个对只问一个问题。真实的应用会 在一次调用中对同一个对问若干个问题。具体做法参见 并行问题 实践手册 和 推测式扇出模式。


接下来

同样的构建模块也出现在 TypeSafe 文档的其他地方:

  • Noul,讲 TypeSafe 如何把一个是/否 问题变成一个得分。
  • 推测式扇出,讲如何 在一次调用中对一份文档提出多个问题。
  • 逐行搜索, 讲另一种按语义而非关键词检索语料库的方式。