Skip to content

Commit 794912d

Browse files
committed
fix: solve P1 bug
1 parent 6887bb3 commit 794912d

2 files changed

Lines changed: 140 additions & 4 deletions

File tree

osa_tool/operations/analysis/paper_claims/claim_extractor.py

Lines changed: 94 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -605,6 +605,20 @@ def _deduplication_batches(self, claims: list[ExtractedClaim]) -> list[list[Extr
605605
)
606606
return batches
607607

608+
def _deduplication_prompt_fits(self, claims: list[ExtractedClaim]) -> bool:
609+
if len(claims) > self.dedup_batch_size:
610+
return False
611+
system = self.prompts.get("paper_claims.deduplication_system")
612+
budget_info = self._input_token_budget(system)
613+
if budget_info is None:
614+
return True
615+
user_budget, encoder = budget_info
616+
try:
617+
return count_tokens(self._deduplication_prompt(claims), encoder) <= user_budget
618+
except Exception as exc:
619+
logger.warning("Token counting failed; assuming deduplication prompt fits current budget: %s", exc)
620+
return True
621+
608622
async def _deduplicate_global_survivors(
609623
self, claims: list[ExtractedClaim]
610624
) -> tuple[list[ExtractedClaim], list[DedupSelection]]:
@@ -672,15 +686,91 @@ async def _deduplicate_claim_group(
672686
return await self._deduplicate_claim_batch(batches[0], request_name=request_name)
673687

674688
logger.info("%s split into %s token-bounded sub-batches", request_name, len(batches))
675-
filtered: list[ExtractedClaim] = []
676-
selections: list[DedupSelection] = []
689+
survivor_batches: list[list[ExtractedClaim]] = []
677690
for batch_index, batch_claims in enumerate(batches, start=1):
678691
batch_filtered, batch_selections = await self._deduplicate_claim_batch(
679692
batch_claims,
680693
request_name=f"{request_name} sub-batch {batch_index}/{len(batches)}",
681694
)
682-
filtered.extend(batch_filtered)
683-
selections.extend(batch_selections)
695+
survivor_batches.append(batch_filtered)
696+
return await self._deduplicate_cross_sub_batch_survivors(survivor_batches, request_name=request_name)
697+
698+
@staticmethod
699+
def _apply_dedup_result(
700+
active: dict[str, ExtractedClaim],
701+
group: list[ExtractedClaim],
702+
kept: list[ExtractedClaim],
703+
) -> None:
704+
kept_by_id = {claim.claim_id: claim for claim in kept}
705+
for claim_id, kept_claim in kept_by_id.items():
706+
current = active.get(claim_id)
707+
if current is None:
708+
continue
709+
active[claim_id] = kept_claim.model_copy(
710+
update={"contradiction": current.contradiction or kept_claim.contradiction}
711+
)
712+
for claim in group:
713+
if claim.claim_id not in kept_by_id:
714+
active.pop(claim.claim_id, None)
715+
716+
async def _deduplicate_cross_sub_batch_survivors(
717+
self,
718+
survivor_batches: list[list[ExtractedClaim]],
719+
*,
720+
request_name: str,
721+
) -> tuple[list[ExtractedClaim], list[DedupSelection]]:
722+
survivors = [claim for batch in survivor_batches for claim in batch]
723+
if len(survivors) <= 1:
724+
return survivors, self._dedup_selections(survivors)
725+
726+
active = {claim.claim_id: claim for claim in survivors}
727+
original_order = {claim.claim_id: index for index, claim in enumerate(survivors)}
728+
total_pairs = len(survivor_batches) * (len(survivor_batches) - 1) // 2
729+
pair_number = 0
730+
for left_index, left_batch in enumerate(survivor_batches):
731+
for right_index in range(left_index + 1, len(survivor_batches)):
732+
pair_number += 1
733+
right_batch = survivor_batches[right_index]
734+
group_ids = [claim.claim_id for claim in [*left_batch, *right_batch]]
735+
group = [active[claim_id] for claim_id in group_ids if claim_id in active]
736+
if len(group) <= 1:
737+
continue
738+
if self._deduplication_prompt_fits(group):
739+
kept, _chosen = await self._deduplicate_claim_batch(
740+
group,
741+
request_name=f"{request_name} cross-sub-batch group {pair_number}/{total_pairs}",
742+
)
743+
self._apply_dedup_result(active, group, kept)
744+
continue
745+
746+
comparison_number = 0
747+
for left_claim in left_batch:
748+
for right_claim in right_batch:
749+
left_active = active.get(left_claim.claim_id)
750+
right_active = active.get(right_claim.claim_id)
751+
if left_active is None or right_active is None:
752+
continue
753+
comparison_number += 1
754+
pair = [left_active, right_active]
755+
if not self._deduplication_prompt_fits(pair):
756+
logger.warning(
757+
"%s: cannot compare verbose claims %s and %s within deduplication input budget",
758+
request_name,
759+
left_active.claim_id,
760+
right_active.claim_id,
761+
)
762+
continue
763+
kept, _chosen = await self._deduplicate_claim_batch(
764+
pair,
765+
request_name=(
766+
f"{request_name} cross-sub-batch pair {pair_number}/{total_pairs}."
767+
f"{comparison_number}"
768+
),
769+
)
770+
self._apply_dedup_result(active, pair, kept)
771+
772+
filtered = sorted(active.values(), key=lambda claim: original_order[claim.claim_id])
773+
selections = self._dedup_selections(filtered)
684774
return filtered, selections
685775

686776
async def _deduplicate_claim_batch(

tests/unit/operations/analysis/paper_claims/test_claim_extractor.py

Lines changed: 46 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -227,6 +227,52 @@ def test_deduplication_batches_respect_model_input_token_budget():
227227
assert all(count_tokens(extractor._deduplication_prompt(batch), "cl100k_base") <= user_budget for batch in batches)
228228

229229

230+
@pytest.mark.asyncio
231+
async def test_deduplicate_claim_group_compares_survivors_across_token_split_sub_batches():
232+
system = PromptLoader().get("paper_claims.deduplication_system")
233+
max_tokens = 80
234+
user_budget = 130
235+
claims = [
236+
extracted_claim("c0001", "Unique claim " + "token " * 30),
237+
extracted_claim("c0002", "Duplicate verbose claim " + "token " * 30),
238+
extracted_claim("c0003", "Duplicate verbose claim " + "token " * 30),
239+
]
240+
handler = FakeHandler(
241+
[
242+
json.dumps(
243+
[
244+
{"claim_id": "c0001", "claim": claims[0].claim, "contradiction": False},
245+
{"claim_id": "c0002", "claim": claims[1].claim, "contradiction": False},
246+
]
247+
),
248+
json.dumps([{"claim_id": "c0003", "claim": claims[2].claim, "contradiction": False}]),
249+
json.dumps(
250+
[
251+
{"claim_id": "c0001", "claim": claims[0].claim, "contradiction": False},
252+
{"claim_id": "c0003", "claim": claims[2].claim, "contradiction": False},
253+
]
254+
),
255+
json.dumps([{"claim_id": "c0002", "claim": claims[1].claim, "contradiction": False}]),
256+
]
257+
)
258+
handler.model_settings = SimpleNamespace(
259+
context_window=count_tokens(system, "cl100k_base") + max_tokens + 256 + user_budget,
260+
max_tokens=max_tokens,
261+
encoder="cl100k_base",
262+
)
263+
264+
filtered, selections = await ClaimExtractor(handler, dedup_batch_size=100)._deduplicate_claim_group(
265+
claims,
266+
request_name="Test deduplication group",
267+
)
268+
269+
assert [claim.claim_id for claim in filtered] == ["c0001", "c0002"]
270+
assert [selection.claim_id for selection in selections] == ["c0001", "c0002"]
271+
assert len(handler.prompts) == 4
272+
assert '"claim_id": "c0002"' in handler.prompts[-1]
273+
assert '"claim_id": "c0003"' in handler.prompts[-1]
274+
275+
230276
@pytest.mark.asyncio
231277
async def test_extract_repairs_invalid_source_text_and_deduplicates():
232278
valid_claim = {

0 commit comments

Comments
 (0)