@@ -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 (
0 commit comments