Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
167 changes: 167 additions & 0 deletions src/sentinel/sentinel_local_index.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,38 @@ def _corpus_if_aligned(
return corpus


def _split_generators(
seed: Optional[int],
) -> Tuple[Optional[torch.Generator], Optional[torch.Generator]]:
"""Derive one independent generator per index side, or (None, None) when unseeded.

A single shared generator would couple the two sides: selecting rows advances it,
so whether the positives drew at all would decide which negatives came out. Keeping
every positive draws nothing, so a grid cell at full positive size would select
different negatives from one that subsampled, and two cells meant to differ along
one axis would quietly differ along both.

Args:
seed: The caller's seed, or None to leave selection unseeded.

Returns:
Tuple of (positive generator, negative generator), both None when seed is None.
"""
if seed is None:
return None, None
# Draw the two seeds from a root generator rather than offsetting the caller's seed
# by hand, which keeps them independent without inventing arithmetic that could
# collide across nearby seeds.
root = torch.Generator().manual_seed(seed)
positive_seed, negative_seed = torch.randint(
high=2**62, size=(2,), generator=root
).tolist()
return (
torch.Generator().manual_seed(positive_seed),
torch.Generator().manual_seed(negative_seed),
)


def _take_rows(
embeddings: torch.Tensor,
corpus: Optional[List[str]],
Expand Down Expand Up @@ -397,6 +429,141 @@ def _apply_negative_ratio(
self.negative_embeddings.shape[0],
)

def _select_subset(
self,
embeddings: torch.Tensor,
corpus: Optional[List[str]],
n_keep: Optional[int],
generator: Optional[torch.Generator],
label: str,
) -> Tuple[torch.Tensor, Optional[List[str]]]:
"""Randomly keep n_keep rows of one side of the index, corpus included.

Args:
embeddings: The embeddings to select from.
corpus: Matching texts, or None.
n_keep: How many rows to keep. None or a value at least as large as the
available rows keeps everything.
generator: Optional seeded generator for a reproducible choice.
label: "positive" or "negative", used in log messages.

Returns:
Tuple of (embeddings, corpus) for the kept rows.
"""
available = embeddings.shape[0]

if n_keep is None or n_keep >= available:
if n_keep is not None and n_keep > available:
LOG.info(
"Requested %d %s examples but the index only has %d - keeping all of them.",
n_keep,
label,
available,
)
# Copy the corpus list so callers cannot mutate the original through the copy.
return embeddings, (list(corpus) if corpus is not None else None)

indices = torch.randperm(available, generator=generator)[:n_keep]
Comment on lines +455 to +466

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

edge case: if we are doing a grid search with positive to negative ratio, if there is one run in the search where we want to keep all the positive rows, then torch.randperm(available, generator=generator) will not run, which means the stateful generator will be at a different state, and thus the negative embedding set will be different.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Addressed.

# Order does not affect semantic_search, but keeping the original relative
# order makes the result far easier to diff and debug.
indices = torch.sort(indices).values
LOG.info("Keeping %d %s examples out of %d", n_keep, label, available)
return _take_rows(embeddings, corpus, indices)

def subsample(
self,
n_positive: Optional[int] = None,
neg_to_pos_ratio: Optional[float] = None,
seed: Optional[int] = None,
) -> "SentinelLocalIndex":
"""Return a smaller copy of this index, reusing the existing embeddings.

Encoding a sentence produces the same numbers regardless of which index it
ends up in, so a small index is just a large one with rows removed. This turns
"re-encode everything" into "copy the rows you want", which is what makes
sweeping index sizes affordable.

Corpus texts are kept aligned with the embeddings they describe, and the
sentence model is shared with the copy rather than reloaded.

Args:
n_positive: How many positive examples to keep. None keeps all of them.
If larger than the index, everything available is kept.
neg_to_pos_ratio: Negatives to keep per kept positive. None leaves the
negatives untouched - note that shrinking the positives alone therefore
*changes* the effective ratio, which is easy to do by accident.
seed: Optional seed making the selection reproducible. Uses private
torch.Generators, so the caller's other randomness is unaffected. Each
side gets its own, which means the negatives chosen for a given seed and
count do not change according to whether the positives were subsampled.

Returns:
A new SentinelLocalIndex. This instance is never modified.

Raises:
ValueError: If the index has no embeddings, or if n_positive or
neg_to_pos_ratio is not positive. Unlike load()'s ratio handling, which
warns and carries on, this raises: a grid search that silently ignored a
bad argument would emit result rows describing an index you did not ask for.
"""
if self.positive_embeddings is None or self.negative_embeddings is None:
raise ValueError(
"Cannot subsample an index without both positive and negative embeddings."
)
if n_positive is not None and n_positive <= 0:
raise ValueError(f"n_positive must be positive, got {n_positive}.")
if neg_to_pos_ratio is not None and neg_to_pos_ratio <= 0:
raise ValueError(
f"neg_to_pos_ratio must be positive, got {neg_to_pos_ratio}."
)

# One generator per side, so neither side's selection depends on how many draws
# the other happened to make. See _split_generators.
positive_generator, negative_generator = _split_generators(seed)

# Positives first: the ratio is defined relative to how many positives survive,
# so that count has to be settled before the negatives can be sized.
positive_embeddings, positive_corpus = self._select_subset(
self.positive_embeddings,
self.positive_corpus,
n_positive,
positive_generator,
"positive",
)

n_negative: Optional[int] = None
if neg_to_pos_ratio is not None:
n_negative = int(positive_embeddings.shape[0] * neg_to_pos_ratio)
if n_negative <= 0:
LOG.info(
"Ratio %.4f against %d positives rounds to zero negatives - keeping 1, "
"since an index with no negatives cannot score anything.",
neg_to_pos_ratio,
positive_embeddings.shape[0],
)
n_negative = 1

negative_embeddings, negative_corpus = self._select_subset(
self.negative_embeddings,
self.negative_corpus,
n_negative,
negative_generator,
"negative",
)

# scale_fn, encoding_kwargs and model_card all carry over. A dropped scale_fn
# would silently change scores for models like E5.
return type(self)(
sentence_model=self.sentence_model,
positive_embeddings=positive_embeddings,
negative_embeddings=negative_embeddings,
scale_fn=self.scale_fn,
encoding_additional_kwargs=self.encoding_kwargs,
positive_corpus=positive_corpus,
negative_corpus=negative_corpus,
model_card=self.model_card,
)

def calculate_rare_class_affinity(
self,
text_samples: List[str],
Expand Down
197 changes: 197 additions & 0 deletions tests/test_sentinel_local_index.py
Original file line number Diff line number Diff line change
Expand Up @@ -637,3 +637,200 @@ def test_misaligned_corpus_is_dropped_with_warning(self, caplog):

assert index.negative_corpus is None
assert "does not match embedding rows" in caplog.text


def _labelled_pair_index(n_positive=20, n_negative=40, dim=8):
"""Index where BOTH sides carry the row-number-in-the-embedding trick."""
return SentinelLocalIndex(
sentence_model=None,
positive_embeddings=torch.arange(n_positive, dtype=torch.float32).repeat(dim, 1).T,
negative_embeddings=torch.arange(n_negative, dtype=torch.float32).repeat(dim, 1).T,
positive_corpus=[f"pos {i}" for i in range(n_positive)],
negative_corpus=[f"neg {i}" for i in range(n_negative)],
scale_fn=lambda s: s * 2,
model_card={"version": "1.0"},
)


def _assert_aligned(embeddings, corpus, prefix):
"""Every row must still sit beside the text it was built from."""
for row, text in zip(embeddings, corpus):
assert text == f"{prefix} {int(row[0].item())}"


class TestSubsample:
"""Resizing an index without re-encoding."""

def test_shapes_for_both_axes(self):
"""Both axes resize as asked, independently and together."""
index = _labelled_pair_index()

only_positive = index.subsample(n_positive=5)
assert only_positive.positive_embeddings.shape[0] == 5
assert only_positive.negative_embeddings.shape[0] == 40 # untouched

only_ratio = index.subsample(neg_to_pos_ratio=0.5)
assert only_ratio.positive_embeddings.shape[0] == 20 # untouched
assert only_ratio.negative_embeddings.shape[0] == 10 # 20 * 0.5

both = index.subsample(n_positive=6, neg_to_pos_ratio=2.0)
assert both.positive_embeddings.shape[0] == 6
assert both.negative_embeddings.shape[0] == 12 # 6 * 2.0

def test_original_index_is_untouched(self):
"""The most important guarantee: subsample() never mutates the receiver.

A grid search loops over sizes. If each call shrank the original, run 2 would
start from run 1's leftovers and every result after the first would be wrong.
"""
index = _labelled_pair_index()
before_positive = index.positive_embeddings.clone()
before_negative = index.negative_embeddings.clone()
before_pos_corpus = list(index.positive_corpus)

for _ in range(3):
index.subsample(n_positive=5, neg_to_pos_ratio=1.0, seed=1)

assert torch.equal(index.positive_embeddings, before_positive)
assert torch.equal(index.negative_embeddings, before_negative)
assert index.positive_corpus == before_pos_corpus
assert index.negative_embeddings.shape[0] == 40

def test_repeated_calls_are_independent(self):
"""Chained calls all measure the same starting index, not each other's output."""
index = _labelled_pair_index()
sizes = [index.subsample(n_positive=n).positive_embeddings.shape[0] for n in (10, 8, 3)]
assert sizes == [10, 8, 3]

def test_alignment_preserved_on_both_sides(self):
"""Corpus texts follow their own embedding rows through the resize."""
index = _labelled_pair_index()
smaller = index.subsample(n_positive=7, neg_to_pos_ratio=2.0, seed=3)

assert len(smaller.positive_corpus) == 7
assert len(smaller.negative_corpus) == 14
_assert_aligned(smaller.positive_embeddings, smaller.positive_corpus, "pos")
_assert_aligned(smaller.negative_embeddings, smaller.negative_corpus, "neg")

def test_same_seed_reproducible_different_seed_not(self):
"""Seeding controls the selection, as everywhere else in the library."""
index = _labelled_pair_index(n_positive=200, n_negative=400)

a = index.subsample(n_positive=50, neg_to_pos_ratio=1.0, seed=42)
b = index.subsample(n_positive=50, neg_to_pos_ratio=1.0, seed=42)
c = index.subsample(n_positive=50, neg_to_pos_ratio=1.0, seed=99)

assert torch.equal(a.positive_embeddings, b.positive_embeddings)
assert a.positive_corpus == b.positive_corpus
assert not torch.equal(a.positive_embeddings, c.positive_embeddings)

def test_negatives_do_not_depend_on_the_positive_code_path(self):
"""The same seed and negative count must select the same negatives either way.

Keeping every positive returns early without drawing, while subsampling them
calls randperm. On a single shared generator that difference leaves the
generator in a different state, so the negatives would silently differ between
two grid cells that were only meant to differ in positive size.

Both calls below land on 20 negatives by different routes: 10 positives at
ratio 2.0 (no positive draw) and 5 positives at ratio 4.0 (a positive draw).
"""
index = _labelled_pair_index(n_positive=10, n_negative=200)

kept_all_positives = index.subsample(neg_to_pos_ratio=2.0, seed=42)
subsampled_positives = index.subsample(
n_positive=5, neg_to_pos_ratio=4.0, seed=42
)

assert kept_all_positives.negative_embeddings.shape[0] == 20
assert subsampled_positives.negative_embeddings.shape[0] == 20
assert torch.equal(
kept_all_positives.negative_embeddings,
subsampled_positives.negative_embeddings,
)
assert (
kept_all_positives.negative_corpus == subsampled_positives.negative_corpus
)

def test_requesting_more_than_available_keeps_everything(self):
"""Over-asking clips rather than erroring, matching load()'s existing behaviour."""
index = _labelled_pair_index(n_positive=20, n_negative=40)
bigger = index.subsample(n_positive=999, neg_to_pos_ratio=999.0)

assert bigger.positive_embeddings.shape[0] == 20
assert bigger.negative_embeddings.shape[0] == 40

def test_no_arguments_returns_equivalent_copy(self):
"""Both arguments None yields an equal but distinct index."""
index = _labelled_pair_index()
copy = index.subsample()

assert copy is not index
assert torch.equal(copy.positive_embeddings, index.positive_embeddings)
assert copy.positive_corpus == index.positive_corpus
# Mutating the copy's corpus list must not reach back into the original.
copy.positive_corpus.append("extra")
assert len(index.positive_corpus) == 20

def test_metadata_carries_over(self):
"""scale_fn, encoding kwargs, model card and the model itself come along.

A dropped scale_fn silently changes scores for models like E5, so this is a
correctness check rather than a tidiness one.
"""
index = _labelled_pair_index()
smaller = index.subsample(n_positive=4)

assert smaller.scale_fn is index.scale_fn
assert smaller.scale_fn(3) == 6
assert smaller.model_card == index.model_card
assert smaller.encoding_kwargs == index.encoding_kwargs
assert smaller.encoding_kwargs["normalize_embeddings"] is True
# The model is large and read-only, so it is shared rather than reloaded.
assert smaller.sentence_model is index.sentence_model

def test_index_without_corpus(self):
"""An index carrying no corpus subsamples fine and stays corpus-free."""
index = SentinelLocalIndex(
sentence_model=None,
positive_embeddings=torch.rand(20, 8),
negative_embeddings=torch.rand(40, 8),
)
smaller = index.subsample(n_positive=5, neg_to_pos_ratio=2.0, seed=1)

assert smaller.positive_embeddings.shape[0] == 5
assert smaller.negative_embeddings.shape[0] == 10
assert smaller.positive_corpus is None
assert smaller.negative_corpus is None

def test_tiny_ratio_keeps_at_least_one_negative(self):
"""A ratio that rounds to zero negatives is clipped to one, not left empty."""
index = _labelled_pair_index(n_positive=20, n_negative=40)
smaller = index.subsample(n_positive=2, neg_to_pos_ratio=0.1) # 2 * 0.1 -> 0

assert smaller.negative_embeddings.shape[0] == 1

@pytest.mark.parametrize(
"kwargs,message",
[
({"n_positive": 0}, "n_positive must be positive"),
({"n_positive": -5}, "n_positive must be positive"),
({"neg_to_pos_ratio": 0}, "neg_to_pos_ratio must be positive"),
({"neg_to_pos_ratio": -1.0}, "neg_to_pos_ratio must be positive"),
],
)
def test_invalid_arguments_raise(self, kwargs, message):
"""Bad sizes raise instead of being quietly ignored.

Silently ignoring them would produce grid-search rows describing an index the
caller never asked for.
"""
index = _labelled_pair_index()
with pytest.raises(ValueError, match=message):
index.subsample(**kwargs)

def test_missing_embeddings_raise(self):
"""Subsampling an index with nothing in it is an error, not a silent no-op."""
index = SentinelLocalIndex(sentence_model=None)
with pytest.raises(ValueError, match="without both positive and negative"):
index.subsample(n_positive=1)
Loading