diff --git a/CHANGELOG.md b/CHANGELOG.md index 00c89253..6933cd62 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,13 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added +- **Configurable LLM Retry Logic**: + - Exposed `max_retries` parameter in `NERExtractor`, `RelationExtractor`, `TripletExtractor` and low-level extraction methods (`extract_entities_llm`, `extract_relations_llm`, `extract_triplets_llm`). + - Defaults to 3 retries to prevent infinite loops during JSON validation failures or API timeouts. + - Propagated retry configuration through chunked processing helpers to ensure consistent behavior for long documents. + - Updated `03_Earnings_Call_Analysis.ipynb` to use `max_retries=3` by default. + ### Added - **Bring Your Own Model (BYOM) Support**: - Enabled full support for custom Hugging Face models in `NERExtractor`, `RelationExtractor`, and `TripletExtractor`. @@ -24,6 +31,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Implemented post-processing logic to clean and validate generated triplets. ### Fixed +- **LLM Extraction Stability**: + - Fixed infinite retry loops in `BaseProvider` by strictly enforcing `max_retries` limit during structured output generation. + - Resolved stuck execution in earnings call analysis notebooks when using smaller models (e.g., Llama 3 8B) that frequently produce invalid JSON. - **Model Parameter Precedence**: - Fixed issue where configuration defaults took precedence over runtime arguments in Hugging Face extractors. Runtime options now correctly override config values. - **Import Handling**: diff --git a/cookbook/use_cases/finance/03_Earnings_Call_Analysis.ipynb b/cookbook/use_cases/finance/03_Earnings_Call_Analysis.ipynb index 44182be3..55bf86dc 100644 --- a/cookbook/use_cases/finance/03_Earnings_Call_Analysis.ipynb +++ b/cookbook/use_cases/finance/03_Earnings_Call_Analysis.ipynb @@ -288,6 +288,7 @@ " llm_model=\"llama-3.1-8b-instant\",\n", " temperature=0.0,\n", " api_key=GROQ_API_KEY,\n", + " max_retries=3,\n", ")\n", "\n", "ENTITY_TYPES = [\"ORGANIZATION\", \"PERSON\", \"MONEY\", \"PERCENT\", \"DATE\", \"EVENT\"]\n", diff --git a/semantica/semantic_extract/methods.py b/semantica/semantic_extract/methods.py index 458dea89..2418bb0a 100644 --- a/semantica/semantic_extract/methods.py +++ b/semantica/semantic_extract/methods.py @@ -1650,6 +1650,7 @@ def extract_relations_llm( silent_fail: bool = False, max_text_length: Optional[int] = None, structured_output_mode: str = "typed", + max_retries: int = 3, **kwargs, ) -> List[Relation]: """ @@ -1662,6 +1663,7 @@ def extract_relations_llm( model: LLM model silent_fail: If True, return empty list on error. If False (default), raise exception. max_text_length: Maximum text length before auto-chunking. None = provider default. + max_retries: Maximum number of retries for LLM calls (default: 3) **kwargs: Additional options """ # Support llm_model parameter to disambiguate from ML model @@ -1674,6 +1676,7 @@ def extract_relations_llm( "model": model, "max_text_length": max_text_length, "structured_output_mode": structured_output_mode, + "max_retries": max_retries, "relation_types": kwargs.get("relation_types"), # Include entities hash/str in cache key implicitly via **cache_params "entities_hash": hash(tuple(sorted([e.text for e in entities]))) if entities else 0 @@ -1745,6 +1748,7 @@ def extract_relations_llm( return _extract_relations_chunked( text, entities, provider=provider, model=model, silent_fail=silent_fail, max_text_length=max_text_length, + max_retries=max_retries, **kwargs ) @@ -1816,6 +1820,8 @@ Entities found in text: {entities_str}""" call_kwargs["temperature"] = kwargs["temperature"] if "verbose" in kwargs: call_kwargs["verbose"] = kwargs["verbose"] + + call_kwargs["max_retries"] = max_retries result_obj = llm.generate_typed(prompt, schema=RelationsResponse, **call_kwargs) if verbose_mode: @@ -1962,6 +1968,7 @@ def _extract_relations_chunked( silent_fail: bool, max_text_length: int, structured_output_mode: str = "typed", + max_retries: int = 3, **kwargs ) -> List[Relation]: """Internal helper to extract relations from long text by chunking.""" @@ -2004,6 +2011,7 @@ def _extract_relations_chunked( silent_fail=False, max_text_length=len(chunk.text) + 1, structured_output_mode=structured_output_mode, + max_retries=max_retries, **limited_kwargs ) future_to_chunk[future] = i @@ -2176,6 +2184,7 @@ def extract_triplets_llm( silent_fail: bool = False, max_text_length: Optional[int] = None, structured_output_mode: str = "typed", + max_retries: int = 3, **kwargs, ) -> List[Triplet]: """ @@ -2189,6 +2198,7 @@ def extract_triplets_llm( model: LLM model silent_fail: If True, return empty list on error. If False (default), raise exception. max_text_length: Maximum text length before auto-chunking. None = provider default. + max_retries: Maximum number of retries for LLM calls (default: 3) **kwargs: Additional options """ # Support llm_model parameter to disambiguate from ML model @@ -2201,6 +2211,7 @@ def extract_triplets_llm( "model": model, "max_text_length": max_text_length, "structured_output_mode": structured_output_mode, + "max_retries": max_retries, "triplet_types": kwargs.get("triplet_types"), # Include entities/relations hash in cache key implicitly via **cache_params "entities_hash": hash(tuple(sorted([e.text for e in entities]))) if entities else 0, @@ -2266,6 +2277,7 @@ def extract_triplets_llm( return _extract_triplets_chunked( text, provider=provider, model=model, silent_fail=silent_fail, max_text_length=max_text_length, + max_retries=max_retries, **kwargs ) @@ -2308,7 +2320,9 @@ Text to extract from: try: # Use typed generation with Pydantic schema - result_obj = llm.generate_typed(prompt, schema=TripletsResponse, **kwargs) + call_kwargs = kwargs.copy() + call_kwargs["max_retries"] = max_retries + result_obj = llm.generate_typed(prompt, schema=TripletsResponse, **call_kwargs) # Convert back to internal Triplet format triplets = [] diff --git a/tests/semantic_extract/test_retry_logic.py b/tests/semantic_extract/test_retry_logic.py new file mode 100644 index 00000000..7c84fe07 --- /dev/null +++ b/tests/semantic_extract/test_retry_logic.py @@ -0,0 +1,264 @@ + +import unittest +from unittest.mock import MagicMock, patch, ANY +import sys +import os +from typing import List, Optional +from pydantic import BaseModel +import importlib.util + +# Add project root to path +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../.."))) + +# Mock external dependencies with __spec__ for importlib checks +mock_spacy = MagicMock() +mock_spacy.__spec__ = MagicMock() +sys.modules["spacy"] = mock_spacy + +sys.modules["instructor"] = MagicMock() +sys.modules["groq"] = MagicMock() + +# Better mock for openai +mock_openai = MagicMock() +mock_openai.__spec__ = MagicMock() +sys.modules["openai"] = mock_openai + +# Mock sentence_transformers and transformers to avoid heavy imports and dependency checks +sys.modules["sentence_transformers"] = MagicMock() +mock_transformers = MagicMock() +mock_transformers.__spec__ = MagicMock() +sys.modules["transformers"] = mock_transformers + +from semantica.semantic_extract import NERExtractor +from semantica.semantic_extract.methods import extract_entities_llm, _extract_entities_chunked, extract_relations_llm, extract_triplets_llm +from semantica.semantic_extract.providers import BaseProvider + +class EntitiesResponse(BaseModel): + entities: List[dict] + +class TestRetryLogic(unittest.TestCase): + + def setUp(self): + self.mock_provider = MagicMock() + self.mock_provider.is_available.return_value = True + self.mock_provider.generate_typed.return_value = MagicMock(entities=[]) + + def test_ner_extractor_init_default(self): + """Test default max_retries in NERExtractor""" + ner = NERExtractor(method="llm", provider="test") + # Check internal config, max_retries not in config means default behavior downstream + self.assertIsNone(ner.config.get("max_retries")) + + def test_ner_extractor_init_custom(self): + """Test custom max_retries in NERExtractor init""" + ner = NERExtractor(method="llm", provider="test", max_retries=5) + self.assertEqual(ner.config.get("max_retries"), 5) + + @patch('semantica.semantic_extract.methods.create_provider') + def test_extract_entities_uses_init_value(self, mock_create_provider): + """Test extract_entities uses initialized max_retries""" + mock_create_provider.return_value = self.mock_provider + + ner = NERExtractor(method="llm", provider="test", max_retries=5) + ner.extract_entities("test text") + + # Verify generate_typed called with max_retries=5 + args, kwargs = self.mock_provider.generate_typed.call_args + self.assertEqual(kwargs.get("max_retries"), 5) + + @patch('semantica.semantic_extract.methods.create_provider') + def test_extract_entities_override(self, mock_create_provider): + """Test extract_entities override max_retries""" + mock_create_provider.return_value = self.mock_provider + + ner = NERExtractor(method="llm", provider="test", max_retries=5) + # Override with 1 + ner.extract_entities("test text", max_retries=1) + + args, kwargs = self.mock_provider.generate_typed.call_args + self.assertEqual(kwargs.get("max_retries"), 1) + + @patch('semantica.semantic_extract.methods.create_provider') + def test_chunked_extraction_propagation(self, mock_create_provider): + """Test max_retries propagation in chunked extraction""" + mock_create_provider.return_value = self.mock_provider + + # Patch TextSplitter where it lives + with patch('semantica.split.TextSplitter') as MockSplitter: + mock_splitter_instance = MockSplitter.return_value + # Mock split to return 2 chunks + mock_chunk1 = MagicMock() + mock_chunk1.text = "chunk1" + mock_chunk2 = MagicMock() + mock_chunk2.text = "chunk2" + mock_splitter_instance.split.return_value = [mock_chunk1, mock_chunk2] + + # Force chunking by setting max_text_length small + extract_entities_llm( + "very long text", + provider="test", + model="test-model", + max_text_length=10, # Force chunking + max_retries=7, + structured_output_mode="typed" + ) + + # Check if generate_typed was called with max_retries=7 for chunks + # It should be called twice (once for each chunk) + self.assertEqual(self.mock_provider.generate_typed.call_count, 2) + + # Check arguments of the calls + call_args_list = self.mock_provider.generate_typed.call_args_list + for args, kwargs in call_args_list: + self.assertEqual(kwargs.get("max_retries"), 7) + + def test_provider_base_logic(self): + """Test BaseProvider logic for max_retries with manual loop""" + provider = BaseProvider() + provider.client = MagicMock() + provider.logger = MagicMock() + provider.generate_structured = MagicMock(side_effect=Exception("Fail")) + + # Mock instructor failing + with patch('semantica.semantic_extract.providers.instructor') as mock_instructor: + # Make instructor client fail + mock_client = MagicMock() + mock_client.chat.completions.create.side_effect = Exception("Instructor Fail") + mock_instructor.from_provider.return_value = mock_client + mock_instructor.from_openai.return_value = mock_client + + # Run with max_retries=2 + try: + provider.generate_typed("prompt", EntitiesResponse, max_retries=2) + except Exception: + pass + + # Should try manual generation exactly 2 times + self.assertEqual(provider.generate_structured.call_count, 2) + + def test_provider_zero_retries(self): + """Test BaseProvider with max_retries=0""" + provider = BaseProvider() + provider.client = MagicMock() + provider.logger = MagicMock() + provider.generate_structured = MagicMock(side_effect=Exception("Fail")) + + # Mock instructor failing + with patch('semantica.semantic_extract.providers.instructor') as mock_instructor: + mock_client = MagicMock() + mock_client.chat.completions.create.side_effect = Exception("Instructor Fail") + mock_instructor.from_provider.return_value = mock_client + mock_instructor.from_openai.return_value = mock_client + + try: + provider.generate_typed("prompt", EntitiesResponse, max_retries=0) + except Exception: + pass + + # Should NOT try manual generation loop (range(0) is empty) + self.assertEqual(provider.generate_structured.call_count, 0) + + @patch('semantica.semantic_extract.methods.create_provider') + def test_relations_retry_propagation(self, mock_create_provider): + """Test max_retries propagation in relation extraction""" + mock_create_provider.return_value = self.mock_provider + + # Create a mock entity + mock_entity = MagicMock() + mock_entity.text = "entity" + mock_entity.start_char = 0 + mock_entity.end_char = 5 + + extract_relations_llm( + "test text", + entities=[mock_entity], + provider="test", + max_retries=4 + ) + + args, kwargs = self.mock_provider.generate_typed.call_args + self.assertEqual(kwargs.get("max_retries"), 4) + + @patch('semantica.semantic_extract.methods.create_provider') + def test_relations_chunked_propagation(self, mock_create_provider): + """Test max_retries propagation in chunked relation extraction""" + mock_create_provider.return_value = self.mock_provider + + with patch('semantica.split.TextSplitter') as MockSplitter: + mock_splitter_instance = MockSplitter.return_value + mock_chunk1 = MagicMock() + mock_chunk1.text = "chunk1" + mock_chunk1.start_index = 0 + mock_chunk1.end_index = 6 + mock_splitter_instance.split.return_value = [mock_chunk1] + + # Create a mock entity + mock_entity = MagicMock() + mock_entity.text = "entity" + mock_entity.start_char = 0 + mock_entity.end_char = 5 + + extract_relations_llm( + "very long text", + entities=[mock_entity], + provider="test", + max_text_length=10, + max_retries=6 + ) + + # Check call count - should be called for the chunk + # Note: _extract_relations_chunked creates a new future for each chunk + # which calls extract_relations_llm, which calls generate_typed + self.assertEqual(self.mock_provider.generate_typed.call_count, 1) + + args, kwargs = self.mock_provider.generate_typed.call_args + self.assertEqual(kwargs.get("max_retries"), 6) + + @patch('semantica.semantic_extract.methods.create_provider') + def test_triplets_retry_propagation(self, mock_create_provider): + """Test max_retries propagation in triplet extraction""" + mock_create_provider.return_value = self.mock_provider + + extract_triplets_llm( + "test text", + entities=[], + relations=[], + provider="test", + max_retries=7 + ) + + args, kwargs = self.mock_provider.generate_typed.call_args + self.assertEqual(kwargs.get("max_retries"), 7) + + @patch('semantica.semantic_extract.methods.create_provider') + def test_triplets_chunked_propagation(self, mock_create_provider): + """Test max_retries propagation in chunked triplet extraction""" + mock_create_provider.return_value = self.mock_provider + + with patch('semantica.split.TextSplitter') as MockSplitter: + mock_splitter_instance = MockSplitter.return_value + mock_chunk1 = MagicMock() + mock_chunk1.text = "chunk1" + mock_chunk1.start_index = 0 + mock_chunk1.end_index = 6 + mock_splitter_instance.split.return_value = [mock_chunk1] + + # Use max_text_length > 100 to pass the minimum viable chunk size check + # and make text longer than that + extract_triplets_llm( + "very long text " * 20, # length > 101 + entities=[], + relations=[], + provider="test", + max_text_length=101, + max_retries=8 + ) + + # Check call count + self.assertEqual(self.mock_provider.generate_typed.call_count, 1) + + args, kwargs = self.mock_provider.generate_typed.call_args + self.assertEqual(kwargs.get("max_retries"), 8) + +if __name__ == "__main__": + unittest.main()