ARM ONLY, NEVER MERGE: z-scored branch fusion

Hybrid without a reranker keeps each database's vector and FTS branches
apart and compares every candidate by per-branch z-score, summing both
branches for a chunk found in each. Databases are then compared by how
exceptional a hit is for them rather than by raw score, which is not
comparable across indexes.

Targets the measured ceiling on rank-and-score fusion: 73% of candidates at
n=4 and 84% at n=8 tie on both score and rank, so no key built from those
two can separate them and they fall to declaration order under every other
arm. Continuous keys should barely collide.

Implementation from the multi-fusion session; branch depth via
HAIKU_RAG_BRANCH_DEPTH, default 20.

Claude-Session: https://claude.ai/code/session_01WhudUtZm6qqiuv8Y1sbwSc
This commit is contained in:
Yiorgis Gozadinos 2026-08-31 17:00:42 +03:00
parent 33bd0be702
commit f4bfc12293
No known key found for this signature in database
2 changed files with 178 additions and 11 deletions

View file

@ -1,4 +1,6 @@
import base64
import os
import statistics
from collections.abc import Sequence
from typing import TYPE_CHECKING
@ -105,20 +107,24 @@ async def search_sources(
fetch_limit = _fetch_limit(client, query, limit)
query_vector = await _embed_query(selected[0], query, resolved)
text = query if isinstance(query, str) else ""
per_source = await gather_all(
*(
c.chunk_repository.search(
query=text,
limit=fetch_limit,
search_type=resolved,
filter=filter,
query_vector=query_vector,
if resolved == "hybrid" and isinstance(query, str) and client.reranker is None:
# ARM C+D (never merge): z-scored branch fusion.
ranked = await _fuse_branches(selected, text, query_vector, filter, limit)
else:
per_source = await gather_all(
*(
c.chunk_repository.search(
query=text,
limit=fetch_limit,
search_type=resolved,
filter=filter,
query_vector=query_vector,
)
for c in selected
)
for c in selected
)
)
ranked = await _fuse(client, selected, query, per_source, limit)
ranked = await _fuse(client, selected, query, per_source, limit)
results: list[SearchResult] = []
for owner, chunk, score in ranked:
@ -199,6 +205,76 @@ async def _fuse(
return [(client, chunk, score) for score, client, chunk in scored[:limit]]
async def _fuse_branches(
clients: list["HaikuRAG"],
query: str,
query_vector: list[float] | None,
filter: str | None,
limit: int,
) -> list[tuple["HaikuRAG", Chunk, float]]:
"""ARM C+D (never merge): one ranked list from every database's vector and
FTS branches, compared by per-branch z-score.
Each branch's scores are normalized against that branch's own candidate
distribution, so databases are compared by how exceptional a hit is for
them rather than by raw score. A chunk in both branches of its database
sums both z-scores. Ties keep arrival order: client order, vector before
FTS, rank within a branch.
"""
async def branches(client: "HaikuRAG"):
return await gather_all(
client.chunk_repository.search(
query=query,
limit=_BRANCH_DEPTH,
search_type="vector",
filter=filter,
query_vector=query_vector,
),
client.chunk_repository.search(
query=query,
limit=_BRANCH_DEPTH,
search_type="fts",
filter=filter,
),
)
per_client = await gather_all(*(branches(c) for c in clients))
totals: dict[tuple[int, str], float] = {}
seen: dict[tuple[int, str], tuple[HaikuRAG, Chunk]] = {}
for position, (client, (vector, fts)) in enumerate(
zip(clients, per_client, strict=True)
):
for branch in (vector, fts):
zs = _z_scores([score for _, score in branch])
for (chunk, _), z in zip(branch, zs, strict=True):
key = (position, chunk.id or chunk.content)
totals[key] = totals.get(key, 0.0) + z
seen.setdefault(key, (client, chunk))
ranked = sorted(seen, key=lambda key: totals[key], reverse=True)[:limit]
return [(*seen[key], totals[key]) for key in ranked]
def _z_scores(scores: list[float]) -> list[float]:
"""Each score as standard deviations above its own list's mean.
A list too short or too flat to carry a distribution normalizes to zeros:
nothing in it is exceptional.
"""
if len(scores) < 2:
return [0.0] * len(scores)
mean = statistics.fmean(scores)
sd = statistics.pstdev(scores)
if sd == 0:
return [0.0] * len(scores)
return [(score - mean) / sd for score in scores]
# ARM C+D (never merge): candidates fetched per branch, an eval grid knob.
_BRANCH_DEPTH = int(os.environ.get("HAIKU_RAG_BRANCH_DEPTH", "20"))
# Reciprocal rank fusion's smoothing constant, the value the literature uses.
_RRF_K = 60

View file

@ -545,6 +545,97 @@ class TestReciprocalRankFusion:
]
class TestZScoredBranchFusion:
"""ARM C+D (never merge): hybrid without a reranker fuses every database's
vector and FTS branches by per-branch z-score."""
@staticmethod
def _client(vector: list[tuple[Chunk, float]], fts: list[tuple[Chunk, float]]):
from types import SimpleNamespace
class Repo:
async def search(
self, query, limit, search_type, filter=None, query_vector=None
):
return vector if search_type == "vector" else fts
return SimpleNamespace(chunk_repository=Repo())
@staticmethod
def _branch(source: str, scores: list[float]) -> list[tuple[Chunk, float]]:
return [
(Chunk(id=f"{source}{i}", content=f"{source} {i}"), score)
for i, score in enumerate(scores)
]
@pytest.mark.asyncio
async def test_a_spike_beats_a_flat_profile(self):
"""The database whose top hit stands out from its own candidates wins,
whatever the raw magnitudes."""
from haiku.rag.client.search import _fuse_branches
spike = self._client(self._branch("a", [0.9, 0.2, 0.19, 0.18]), [])
flat = self._client(self._branch("b", [5.0, 4.99, 4.98, 4.97]), [])
ranked = await _fuse_branches([flat, spike], "q", None, None, 2)
assert [chunk.id for _, chunk, _ in ranked] == ["a0", "b0"]
@pytest.mark.asyncio
async def test_agreement_within_a_database_sums(self):
"""A chunk topping both of its database's branches carries both
z-scores."""
from haiku.rag.client.search import _fuse_branches
scores = [0.9, 0.2, 0.19, 0.18]
both = self._client(self._branch("a", scores), self._branch("a", scores))
one = self._client(self._branch("b", scores), [])
ranked = await _fuse_branches([one, both], "q", None, None, 3)
assert [chunk.id for _, chunk, _ in ranked][:2] == ["a0", "b0"]
assert ranked[0][2] == pytest.approx(2 * ranked[1][2])
@pytest.mark.asyncio
async def test_a_flat_or_thin_branch_is_nothing_special(self):
from haiku.rag.client.search import _z_scores
assert _z_scores([]) == []
assert _z_scores([0.9]) == [0.0]
assert _z_scores([0.5, 0.5, 0.5]) == [0.0, 0.0, 0.0]
@pytest.mark.asyncio
async def test_hybrid_without_a_reranker_fetches_both_branches(
self, tmp_path, monkeypatch, query_embedding
):
from haiku.rag.store.repositories.chunk import ChunkRepository
config = _config(tmp_path, ["alpha", "beta"])
await _seed(config, "alpha", ["alpha document about cats"])
await _seed(config, "beta", ["beta document about cats"])
monkeypatch.setattr(HaikuRAG, "reranker", property(lambda self: None))
asked: list[tuple[str, int]] = []
search = ChunkRepository.search
async def spy(self, *args, **kwargs):
asked.append((kwargs["search_type"], kwargs["limit"]))
return await search(self, *args, **kwargs)
monkeypatch.setattr(ChunkRepository, "search", spy)
async with HaikuRAG(config=config) as rag:
results = await rag.search("cats", limit=3)
assert sorted(asked) == [
("fts", 20),
("fts", 20),
("vector", 20),
("vector", 20),
]
assert {r.source for r in results} == {"alpha", "beta"}
class TestFusingWhatARerankerReturns:
@pytest.mark.asyncio
async def test_a_reranker_returning_copies_is_named(self, tmp_path):