Compare commits

...
70 Commits
Author SHA1 Message Date
KaifAhmad1 b382a7df6e chore: release version 0.2.4 2026-01-22 12:50:07 +05:30
Mohd Kaif b35081e015 Delete examples/demo_ontology_ingest.py 2026-01-21 18:27:06 +05:30
Mohd Kaif 7459393eea Merge pull request #214 from Hawksight-AI/ontology
feat(ontology): Implement OntologyIngestor and update exports
2026-01-21 13:51:54 +05:30
KaifAhmad1 b96e71ae72 feat(ontology): Implement OntologyIngestor and update exports
- Added OntologyIngestor in semantica/ingest/ontology_ingestor.py
- Updated semantica/ontology/__init__.py to export OntologyIngestor
- Updated semantica/ingest/methods.py to use OntologyIngestor
- Added tests for ontology ingestion
- Cleaned up temporary files
2026-01-21 13:46:46 +05:30
KaifAhmad1 fa8544c6d6 Release v0.2.3: Update version, changelog, and documentation 2026-01-20 12:08:46 +05:30
Mohd Kaif 87649b7422 Merge pull request #213 from Hawksight-AI/docs
Fix earnings call analysis notebook: attribute access and export logic
2026-01-20 01:52:42 +05:30
KaifAhmad1 d91619f191 Fix earnings call analysis notebook: attribute access and export logic 2026-01-20 01:51:29 +05:30
Mohd Kaif 064a0db7e6 Merge pull request #212 from Hawksight-AI/docs
Optimize Vector DB Storage in Earnings Call Analysis Notebook
2026-01-19 16:37:33 +05:30
KaifAhmad1 8214acc675 optimize vector db storage in earnings call analysis 2026-01-19 16:32:22 +05:30
Mohd Kaif 2bf55485ff Merge pull request #211 from Hawksight-AI/vector-store
Vector Store Performance Optimization
2026-01-19 13:40:44 +05:30
KaifAhmad1 1568237ce7 Add high-performance VectorStore ingestion and docs 2026-01-19 13:32:16 +05:30
Mohd Kaif f6c9d50e03 Merge pull request #210 from Hawksight-AI/docs
docs: update earnings call analysis notebook
2026-01-18 23:54:44 +05:30
KaifAhmad1 d9117b7c2f docs: update earnings call analysis notebook 2026-01-18 23:53:07 +05:30
Mohd Kaif 0eabfb861e Merge pull request #209 from Hawksight-AI/kg
Fix GraphBuilder External Relationships (#208, #206)
2026-01-18 22:13:16 +05:30
KaifAhmad1 9f77dfb761 Fix GraphBuilder external relationships; refs #208 #206 2026-01-18 22:10:02 +05:30
Mohd Kaif c990d09bd3 Merge pull request #205 from Hawksight-AI/docs
docs: changelog entry for JupyterLab progress flag (#181)
2026-01-17 17:14:54 +05:30
Mohd Kaif 9ebacf43c3 Update CHANGELOG.md 2026-01-17 17:09:41 +05:30
KaifAhmad1 7958ae78f6 docs: changelog entry for JupyterLab progress flag (#181) 2026-01-17 17:03:51 +05:30
Mohd Kaif 2c61fe6cda Merge pull request #204 from Hawksight-AI/utils
feat: allow disabling Jupyter progress output (#181)
2026-01-17 16:44:03 +05:30
KaifAhmad1 92b850ac26 feat: allow disabling Jupyter progress output (#181) 2026-01-17 16:40:15 +05:30
Mohd Kaif f7bd7016c5 Merge pull request #203 from Hawksight-AI/utils
Circular import between `pipeline_builder` and `pipeline_validator`
2026-01-17 16:12:02 +05:30
KaifAhmad1 8671385cbf fix: break pipeline circular import (#192, #193) and update changelog 2026-01-17 16:02:21 +05:30
Mohd Kaif b358acfabf Merge pull request #202 from Hawksight-AI/docs
Update Coockbook
2026-01-16 23:08:03 +05:30
KaifAhmad1 a39ec5fd20 Faster, class-based dedup: DuplicateDetector+EntityMerger with strict thresholds; build graph from deduplicated outputs; clean prints 2026-01-16 22:32:19 +05:30
KaifAhmad1 bbd6764215 Use deduplicated entities/relationships; optimize and clean deduplication; disable extra merging in GraphBuilder 2026-01-16 18:18:50 +05:30
KaifAhmad1 1b0b0551db Update Earnings Call Analysis notebook 2026-01-16 17:54:44 +05:30
Mohd Kaif a6b102fa3d Merge pull request #201 from don-simpson/feature/amazon-neptune-setup
feat: Added CloudFormation template and cookbook instructions for Amazon Neptune
2026-01-16 12:37:23 +05:30
Don Simpson 65d99f7f8a Added CloudFormation template that creates a dev cluster with a single [t3 instance](https://docs.aws.amazon.com/neptune/latest/userguide/manage-console-instances-t3.html) configured with a [public endpoint](https://docs.aws.amazon.com/neptune/latest/userguide/neptune-public-endpoints.html) and IAM Auth enabled (required for public endpoint), and creates an IAM User using least-privilege principles. See Get started with Neptune Database for free on the [Amazon Neptune pricing page](https://aws.amazon.com/neptune/pricing/).
Includes the CloudFormation template in the same directory as the [Amazon Neptune Cookbook](https://github.com/Hawksight-AI/semantica/blob/main/cookbook/introduction/21_Amazon_Neptune_Store.ipynb) and references it as a prerequisite in the cookbook.
2026-01-15 18:48:49 -05:00
Mohd Kaif 9b81137b26 Merge pull request #200 from Hawksight-AI/docs
Update earnings call analysis notebook with relation extraction fixes
2026-01-16 03:05:07 +05:30
KaifAhmad1 653523efeb Update earnings call analysis notebook with relation extraction fixes
- Update notebook to use corrected RelationExtractor API
- Move provider/model parameters to initialization
- Add verbose logging for debugging
- Include working relation extraction examples
2026-01-16 03:03:11 +05:30
Mohd Kaif ba04421d9b Merge pull request #199 from Hawksight-AI/docs
Update changelog for LLM relation extraction fixes
2026-01-16 02:59:44 +05:30
KaifAhmad1 5d3fe51dbd Update changelog for LLM relation extraction fixes
- Add comprehensive changelog entry for relation extraction parsing fixes
- Document breaking changes and new test coverage
- Update with provider normalization and JSON fallback details
2026-01-16 02:58:28 +05:30
Mohd Kaif f20782f517 Merge pull request #198 from Hawksight-AI/semantic-extract
Fix LLM Relation Extraction
2026-01-16 01:22:37 +05:30
KaifAhmad1 96dc5d754a Fix LLM relation extraction parsing and add tests
- Harden LLM relation extraction result handling to parse instructor/OpenAI/Groq variations
- Add structured JSON fallback when typed generation yields zero relations
- Strip acceptance of extra kwargs like max_tokens/max_entities_prompt in relation extraction internals
- Add comprehensive unit tests with mocked LLM provider
- Add integration tests for Groq provider with environment variable API key
- Ensure relation extraction completes and returns results when model identifies relations
2026-01-16 01:19:33 +05:30
Mohd Kaif cf84526cc7 Merge pull request #197 from Hawksight-AI/semantic-extract
Robust LLM Extraction and Groq 401 Fix
2026-01-15 22:46:20 +05:30
KaifAhmad1 5ad20abeab fix(semantic_extract): fix Groq 401 error and improve LLM provider robustness with instructor.from_provider 2026-01-15 22:43:11 +05:30
Mohd Kaif ade08a65ae Merge pull request #196 from Hawksight-AI/semantic-extract
Enhance RelationExtractor with core fixes and verbose logs
2026-01-15 19:03:36 +05:30
KaifAhmad1 fb25644fa7 Enhance RelationExtractor with core fixes and verbose logs
- Fix excessive entities being passed to LLM in RelationExtractor
- Add comprehensive 'Heartbeat' verbose logs to methods.py and providers.py
- Ensure robust API key handling and explicit error reporting
2026-01-15 19:00:42 +05:30
Mohd Kaif 63899f2427 Merge pull request #195 from Hawksight-AI/semantic-extract
Robust Semantic Extraction - API Key Handling & Error Reporting
2026-01-15 18:02:01 +05:30
KaifAhmad1 fd6e058275 feat(semantic_extract): enhance error reporting and API key robustness 2026-01-15 17:59:04 +05:30
Mohd Kaif 23d8207ef5 Merge pull request #194 from Hawksight-AI/semantic-extract
Robust API Key Handling in Semantic Extract Module
2026-01-15 16:32:52 +05:30
KaifAhmad1 f2a11fc8ad fix: robust api_key handling in semantic_extract module 2026-01-15 16:29:52 +05:30
KaifAhmad1 c6316ba4bd Release 0.2.2 2026-01-15 00:42:07 +05:30
Mohd Kaif b6d630fc74 Merge pull request #191 from Hawksight-AI/semantic-extract
Improve `semantic_extract` performance and add Groq LLM smoke tests
2026-01-14 17:21:26 +05:30
Mohd Kaif 3f2cb49e50 Delete PR_DESCRIPTION.md 2026-01-14 17:17:32 +05:30
KaifAhmad1 c7814616a9 Improve semantic_extract performance and add Groq LLM smoke tests 2026-01-14 17:11:26 +05:30
Mohd Kaif 531014fbda Update version and description in pyproject.toml 2026-01-14 14:05:36 +05:30
Mohd Kaif 1cf9b34e3e Merge pull request #190 from Hawksight-AI/utils
docs: update CHANGELOG.md with recent changes
2026-01-14 12:51:40 +05:30
KaifAhmad1 2e81c86489 docs: update CHANGELOG.md with recent changes 2026-01-14 12:49:29 +05:30
Mohd Kaif 1690fec3f7 Merge pull request #189 from Hawksight-AI/utils
resolve dependencies, migrate Gemini SDK, and sanitize notebooks
2026-01-14 12:42:38 +05:30
KaifAhmad1 72a6ddb48f Merge remote-tracking branch 'origin/utils' into utils 2026-01-14 12:38:44 +05:30
KaifAhmad1 a5da533d55 chore: resolve dependencies, migrate Gemini SDK, and sanitize notebooks 2026-01-14 12:37:29 +05:30
Mohd Kaif be8856cfcf Merge pull request #188 from Hawksight-AI/semantic-extract
[SECURITY] Enhance caching security by excluding sensitive keys and using SHA-256
2026-01-14 00:25:46 +05:30
KaifAhmad1 d2e599bcb0 [SECURITY] Enhance caching security by excluding sensitive keys and using SHA-256 2026-01-14 00:22:41 +05:30
Mohd Kaif 05d0bbf86c Merge pull request #187 from Hawksight-AI/semantic-extract
Performance Bottlenecks and Scaling Limitations in semantic_extract
2026-01-14 00:15:06 +05:30
KaifAhmad1 dd7fcd3ddb [FEATURE] Performance Bottlenecks and Scaling Limitations in semantic_extract #186
- Implemented high-throughput parallel batch processing across all core extractors (NERExtractor, RelationExtractor, TripletExtractor, EventDetector, SemanticNetworkExtractor) using ThreadPoolExecutor.

- Added max_workers configuration parameter (default: 1) to all extractor extract() methods.

- Implemented parallel processing for large document chunking in _extract_entities_chunked and _extract_relations_chunked.

- Enhanced ProgressTracker to be thread-safe.

- Optimized setUpClass in tests to reduce Groq LLM initialization overhead.

- Updated documentation and usage examples.
2026-01-14 00:11:30 +05:30
Mohd Kaif 43f55e1028 Delete RELEASE_NOTES_v0.2.0.md 2026-01-13 00:33:59 +05:30
Mohd Kaif e20c522c62 Merge pull request #185 from Hawksight-AI/docs
Update Earning Call Notebook
2026-01-13 00:13:16 +05:30
KaifAhmad1 fd9f0b2526 Add all changes 2026-01-13 00:10:44 +05:30
Mohd Kaif ccaadf6299 Merge pull request #180 from Hawksight-AI/docs
Release v0.2.1: Stability Fixes
2026-01-12 17:52:43 +05:30
KaifAhmad1 428fc3b83a chore(release): bump version to 0.2.1 and update release docs 2026-01-12 17:48:07 +05:30
Mohd Kaif 09cf3ed132 Merge pull request #179 from Hawksight-AI/docs
Resolve TypeError in Earnings Call Analysis Notebook (#177)
2026-01-12 17:35:36 +05:30
KaifAhmad1 58686d409b fix(cookbook): resolve TypeError in earnings call analysis step 7 #177 2026-01-12 17:32:16 +05:30
Mohd Kaif 6d5fbc8b63 Merge pull request #178 from Hawksight-AI/semantic-extract
Resolve Incomplete Output (#176), Relax Constraints, and Add Groq Support
2026-01-12 17:18:54 +05:30
KaifAhmad1 8c3f7f1f0a fix(semantic-extract): resolve incomplete output #176, relax constraints, and add Groq support 2026-01-12 17:15:21 +05:30
Mohd Kaif 4acad23a4d Merge pull request #174 from Hawksight-AI/docs
Update Earnings Call Analysis Notebook (Finance Use Case)
2026-01-11 23:31:13 +05:30
KaifAhmad1 cd1435ee10 Save changes to Earnings Call Analysis notebook 2026-01-11 23:25:28 +05:30
KaifAhmad1 68f0a1d4d9 docs: Update PyPI version badge to shields.io 2026-01-10 23:44:35 +05:30
KaifAhmad1 a47274593b docs: Add v0.2.0 release notes 2026-01-10 23:36:51 +05:30
KaifAhmad1 d8e04c29e9 Security fix: Upgrade protobuf to 4.25.8 and add PR description 2026-01-07 19:11:58 +05:30
65 changed files with 6799 additions and 1677 deletions
+8 -1
View File
@@ -5,6 +5,7 @@ repos:
- id: trailing-whitespace
- id: end-of-file-fixer
- id: check-yaml
exclude: 'neptune-setup\.yaml$'
- id: check-json
- id: check-toml
- id: check-added-large-files
@@ -49,9 +50,15 @@ repos:
hooks:
- id: yamllint
args: ['-d', '{extends: default, rules: {line-length: {max: 120}}}']
exclude: 'neptune-setup\.yaml$'
- repo: https://github.com/aws-cloudformation/cfn-lint
rev: v1.43.3
hooks:
- id: cfn-lint
files: 'neptune-setup\.yaml$'
# Removed slow hooks for faster development:
# - mypy: Type checking (can be run manually or in CI)
# - bandit: Security scanning (can be run separately)
# - pytest: Testing (should be run manually, not on every commit)
+135
View File
@@ -7,6 +7,141 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased]
## [0.2.4] - 2026-01-22
### Added
- **Ontology Ingestion Module**:
- Implemented `OntologyIngestor` in `semantica.ingest` for parsing RDF/OWL files (Turtle, RDF/XML, JSON-LD, N3) into standardized `OntologyData` objects.
- Added `ingest_ontology` convenience function and integrated it into the unified `ingest(source_type="ontology")` interface.
- Added recursive directory scanning support for batch ontology ingestion.
- Exposed ingestion tools in `semantica.ontology` for better discoverability.
- Added `OntologyData` dataclass for consistent metadata handling (source path, format, timestamps).
- **Documentation**:
- **Ontology Usage Guide**: Updated `ontology_usage.md` with comprehensive examples for single-file and directory ingestion.
- **API Reference**: Updated `ontology.md` with `OntologyIngestor` class documentation and method details.
- **Tests**:
- **Comprehensive Test Suite**: Added `tests/ingest/test_ontology_ingestor.py` covering all supported formats, error handling, and unified interface integration.
- **Demo Script**: Added `examples/demo_ontology_ingest.py` for end-to-end usage demonstration.
## [0.2.3] - 2026-01-20
### Fixed
- **LLM Relation Extraction Parsing**:
- Fixed relation extraction returning zero relations despite successful API calls to Groq and other providers
- Normalized typed responses from instructor/OpenAI/Groq to consistent dict format before parsing
- Added structured JSON fallback when typed generation yields zero relations to avoid silent empty outputs
- Removed acceptance of extra kwargs (`max_tokens`, `max_entities_prompt`) from relation extraction internals
- Filtered kwargs passed to provider LLM calls to only `temperature` and `verbose`
- **API Parameter Handling**:
- Limited kwargs forwarded in chunked extraction helper to prevent parameter leakage
- Ensured minimal, safe parameters are passed to provider calls
- **Pipeline Circular Import (Issues #192, #193)**:
- Fixed circular import between `pipeline_builder` and `pipeline_validator` triggered during `semantica.pipeline` import
- Lazy-loaded `PipelineValidator` inside `PipelineBuilder.__init__` and guarded type hints with `TYPE_CHECKING`
- Ensured `from semantica.deduplication import DuplicateDetector` no longer fails even when pipeline module is imported
- **JupyterLab Progress Output (Issue #181)**:
- Added `SEMANTICA_DISABLE_JUPYTER_PROGRESS` environment variable to disable rich Jupyter/Colab progress tables
- When enabled, progress falls back to console-style output, preventing infinite scrolling and JupyterLab out-of-memory errors
### Added
- **Comprehensive Test Suite**:
- - Added unit tests (`tests/test_relations_llm.py`) with mocked LLM provider covering both typed and structured response paths
- - Added integration tests (`tests/integration/test_relations_groq.py`) for real Groq API calls with environment variable API key
- - Tests validate relation extraction completion and result parsing across different response formats
- **Amazon Neptune Dev Environment**:
- - Added CloudFormation template (`cookbook/introduction/neptune-setup.yaml`) to provision a dev Neptune cluster with public endpoint and IAM auth enabled
- - Documented deployment, cost estimates, and IAM User vs IAM Role best practices in `cookbook/introduction/21_Amazon_Neptune_Store.ipynb`
- - Added `cfn-lint` to `.pre-commit-config.yaml` for validating CloudFormation templates while excluding `neptune-setup.yaml` from generic YAML linters
- **Vector Store High-Performance Ingestion**:
- - Added `VectorStore.add_documents` for high-throughput ingestion with automatic embedding generation, batching, and parallel processing
- - Added `VectorStore.embed_batch` helper for generating embeddings for lists of texts without immediately storing them
- - Enabled default parallel ingestion in `VectorStore` with `max_workers=6` for common workloads
- - Added dedicated documentation page `docs/vector_store_usage.md` describing high-performance vector store usage and configuration
- - Added `tests/vector_store/test_vector_store_parallel.py` covering parallel vs sequential performance, error handling, and edge cases for `add_documents` and `embed_batch`
### Changed
- **Relation Extraction API**:
- - Simplified parameter interface by removing unused kwargs that were previously ignored
- - Improved error handling and verbose logging for debugging relation extraction issues
- - Enhanced robustness of post-response parsing across different LLM providers
- **Vector Store Defaults and Examples**:
- - Standardized `VectorStore` default concurrency to `max_workers=6` for parallel ingestion
- - Updated vector store reference documentation and usage guides to rely on implicit defaults instead of requiring manual `max_workers` configuration in examples
## [0.2.2] - 2026-01-15
### Added
- **Parallel Extraction Engine**:
- Implemented high-throughput parallel batch processing across all core extractors (`NERExtractor`, `RelationExtractor`, `TripletExtractor`, `EventDetector`, `SemanticNetworkExtractor`) using `concurrent.futures.ThreadPoolExecutor`.
- Added `max_workers` configuration parameter (default: 1) to all extractor `extract()` methods, allowing users to tune concurrency based on available CPU cores or API rate limits.
- **Parallel Chunking**: Implemented parallel processing for large document chunking in `_extract_entities_chunked` and `_extract_relations_chunked`, significantly reducing latency for long-form text analysis.
- **Thread-Safe Progress Tracking**: Enhanced `ProgressTracker` to handle concurrent updates from multiple threads without race conditions during batch processing.
- **Semantic Extract Performance & Regression**:
- Added edge-case regression suite covering max worker defaults, LLM prompt entity filtering, and extractor reuse.
- Added a runnable real-use-case benchmark script for batch latency across `NERExtractor`, `RelationExtractor`, `TripletExtractor`, `EventDetector`, `SemanticAnalyzer`, and `SemanticNetworkExtractor`.
- Added Groq LLM smoke tests that exercise LLM-based entities/relations/triplets when `GROQ_API_KEY` is available via environment configuration.
### Security
- **Credential Sanitization**:
- Removed hardcoded API keys from 8 cookbook notebooks to prevent secret leakage.
- Enforced environment variable usage for `GROQ_API_KEY` across all examples.
- **Secure Caching**:
- Updated `ExtractionCache` to exclude sensitive parameters (e.g., `api_key`, `token`, `password`) from cache key generation, preventing secret leakage and enabling safe cache sharing.
- Upgraded cache key hashing algorithm from MD5 to **SHA-256** for enhanced collision resistance and security.
### Changed
- **Gemini SDK Migration**:
- Migrated `GeminiProvider` to use the new `google-genai` SDK (v0.1.0+) to address deprecation warnings.
- Implemented graceful fallback to `google.generativeai` for backward compatibility.
- **Dependency Resolution**:
- Pinned `opentelemetry-api` and `opentelemetry-sdk` to `1.37.0` to resolve pip conflicts.
- Updated `protobuf` and `grpcio` constraints for better stability.
- **Entity Filtering Scope**:
- Removed entity filtering from non-LLM extraction flows to avoid accuracy regressions.
- Applied entity downselection only to LLM relation prompt construction, while matching returned entities against the full original entity list.
- **Batch Concurrency Defaults**:
- Standardized `max_workers` defaulting across `semantic_extract` and tuned for low-latency: ML-backed methods default to single-worker, while pattern/regex/rules/LLM/huggingface methods use a higher parallelism default capped by CPU.
- Raised the global `optimization.max_workers` default to 8 for better throughput on batch workloads.
### Performance
- **Bottleneck Optimization (GitHub Issue #186)**:
- **Resolved Bottleneck #1 (Sequential Processing)**: Replaced sequential `for` loops with parallel execution for both document-level batches and intra-document chunks.
- **Performance Gains**: Achieved **~1.89x speedup** in real-world extraction scenarios (tested with Groq `llama-3.3-70b-versatile` on standard datasets).
- **Initialization Optimization**: Refactored test suite to use class-level `setUpClass` for LLM provider initialization, eliminating redundant API client creation overhead.
- **Low-Latency Entity Matching**:
- Avoided heavyweight embedding stack imports on common matches by improving fast matching heuristics and short-circuiting before embedding similarity.
- Optimized entity matching to prioritize exact/substring/word-boundary matches and only fall back to embedding similarity when needed, reducing CPU overhead in LLM relation/triplet mapping.
## [0.2.1] - 2026-01-12
### Fixed
- **LLM Output Stability (Bug #176)**:
- Fixed incomplete JSON output issues by correctly propagating `max_tokens` parameter in `extract_relations_llm`.
- Implemented automatic error handling that halves chunk sizes and retries when LLM context or output limits are exceeded.
- Fixed `AttributeError` in provider integration by ensuring consistent parameter passing via `**kwargs`.
- **Constraint Relaxations**:
- Removed hardcoded `max_length` constraints from `Entity`, `Relation`, and `Triplet` classes to support long-form semantic extraction (e.g., long descriptions or names).
- Fixed orchestrator lazy property initialization and configuration normalization logic in `Orchestrator`.
- Resolved `AssertionError` in orchestrator tests by aligning test mocks with production component usage.
- Fixed dependency compatibility issues by pinning `protobuf>=5.29.1,<7.0` and `grpcio>=1.71.2`.
- Added missing dependencies `GitPython` and `chardet` to `pyproject.toml`.
- Verified and aligned `FileObject.text` property usage in GraphRAG notebooks for consistent content decoding.
### Changed
- **Chunking Defaults**:
- Increased default `max_text_length` for auto-chunking to **64,000 characters** (from 32k/16k) for OpenAI, Anthropic, Gemini, Groq, and DeepSeek providers.
- Unified chunking logic across `extract_entities_llm`, `extract_relations_llm`, and `extract_triplets_llm`.
- **Groq Support**:
- Standardized Groq provider defaults to use `llama-3.3-70b-versatile` with a 64k context window.
- Added native support for `max_tokens` and `max_completion_tokens` to prevent output truncation.
### Added
- **Testing**:
- Added `tests/reproduce_issue_176.py` to validate `max_tokens` propagation and chunking behavior across all extractors.
## [0.2.0] - 2026-01-10
### Added
+2 -2
View File
@@ -6,7 +6,7 @@
[![Python 3.8+](https://img.shields.io/badge/python-3.8+-blue.svg)](https://www.python.org/downloads/)
[![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://opensource.org/licenses/MIT)
[![PyPI version](https://badge.fury.io/py/semantica.svg)](https://pypi.org/project/semantica/)
[![PyPI version](https://img.shields.io/pypi/v/semantica.svg)](https://pypi.org/project/semantica/)
[![Monthly Downloads](https://img.shields.io/pypi/dm/semantica)](https://pypi.org/project/semantica/)
[![Total Downloads](https://static.pepy.tech/badge/semantica)](https://pepy.tech/project/semantica)
[![Discord](https://img.shields.io/badge/Discord-Join%20Us-7289da?style=flat&logo=discord&logoColor=white)](https://discord.gg/pMHguUzG)
@@ -28,7 +28,7 @@
*The missing fabric between raw data and AI engineering. A comprehensive open-source framework for building semantic layers and knowledge engineering systems that transform unstructured data into AI-ready knowledge — powering Knowledge Graph-Powered RAG (GraphRAG), AI Agents, Multi-Agent Systems, and AI applications with structured semantic knowledge.*
**100% Open Source****MIT Licensed****Latest Version: 0.2.0****Production Ready****Community Driven**
**100% Open Source****MIT Licensed****Latest Version: 0.2.4****Production Ready****Community Driven**
[**Discord**](https://discord.gg/pMHguUzG)
+3 -3
View File
@@ -26,10 +26,10 @@ Before releasing, ensure:
The project uses GitHub Actions for automated releases to PyPI.
1. **Tag the commit**: Create a new git tag for the version (e.g., `v0.2.0`).
1.29. **Tag the commit**: Create a new git tag for the version (e.g., `v0.2.3`).
```bash
git tag -a v0.2.0 -m "Release v0.2.0"
git push origin v0.2.0
git tag -a v0.2.3 -m "Release v0.2.3"
git push origin v0.2.3
```
2. **GitHub Action**: The `Release` workflow will automatically trigger, build the package, create a GitHub Release, and publish to PyPI using Trusted Publishing.
+3
View File
@@ -6,6 +6,9 @@ We actively support the following versions of Semantica with security updates:
| Version | Supported |
| ------- | ------------------ |
| 0.2.3 | :white_check_mark: |
| 0.2.2 | :white_check_mark: |
| 0.2.1 | :white_check_mark: |
| 0.2.0 | :white_check_mark: |
| 0.1.1 | :white_check_mark: |
| 0.1.0 | :white_check_mark: |
@@ -25,6 +25,57 @@
"- AWS credentials configured (boto3, environment variables, or IAM role)\n",
"- Network access to your Neptune cluster (VPC, security groups)\n",
"\n",
"#### Quick Setup with CloudFormation\n",
"\n",
"If you don't have a Neptune cluster, use the provided CloudFormation template to create one with a public endpoint and IAM authentication:\n",
"\n",
"```bash\n",
"# Deploy the Neptune stack (takes ~15-20 minutes)\n",
"aws cloudformation create-stack \\\n",
" --stack-name semantica-neptune \\\n",
" --template-body file://neptune-setup.yaml \\\n",
" --capabilities CAPABILITY_NAMED_IAM\n",
"\n",
"# Wait for stack creation to complete\n",
"aws cloudformation wait stack-create-complete --stack-name semantica-neptune\n",
"\n",
"# Get the outputs (endpoint, port, credentials)\n",
"aws cloudformation describe-stacks --stack-name semantica-neptune \\\n",
" --query 'Stacks[0].Outputs' --output table\n",
"```\n",
"\n",
"The template creates:\n",
"- VPC with public subnets and Internet Gateway\n",
"- Neptune cluster (`db.t3.medium`) with IAM authentication enabled\n",
"- IAM user with least-privilege access for OpenCypher queries\n",
"- Security group allowing Bolt protocol (port 8182) access\n",
"\n",
"> ⚠️ **Security Note**: This template creates an IAM User with static access keys for simplicity in demo/test environments. For production use, we recommend IAM Roles (EC2 instance roles, ECS task roles, Lambda execution roles) which provide temporary credentials that are automatically rotated. The secret access key in the Cloudformation outputs is provided in plaintext to simplify initial setup - in production, use AWS Secrets Manager.\n",
"\n",
"**Outputs:**\n",
"- `NeptuneEndpoint` - Cluster hostname (use as `NEPTUNE_ENDPOINT`)\n",
"- `NeptunePort` - 8182 (use as `NEPTUNE_PORT`)\n",
"- `AwsAccessKeyId` - IAM user access key (use as `AWS_ACCESS_KEY_ID`)\n",
"- `AwsSecretAccessKey` - IAM user secret key in **plaintext** (use as `AWS_SECRET_ACCESS_KEY`)\n",
"- `AwsRegion` - Deployment region (use as `AWS_REGION`)\n",
"\n",
"**Cleanup:**\n",
"```bash\n",
"aws cloudformation delete-stack --stack-name semantica-neptune\n",
"```\n",
"\n",
"**Estimated Monthly Cost (approximately 100-105 USD/month at 100% utilization):**\n",
"\n",
"| Resource | Cost (USD) |\n",
"| --- | --- |\n",
"| Neptune db.t3.medium instance | ~96/month (0.132/hr) |\n",
"| Storage (10 GB) | ~1/month |\n",
"| I/O requests | ~1-5/month |\n",
"| Public IPv4 address | ~3.60/month (0.005/hr) |\n",
"| VPC, subnets, route tables, Internet Gateway, IAM | No Additional Charge |\n",
"\n",
"> **Free Tier**: New Neptune users get 30 days free (750 hours of db.t3.medium, 10M I/Os, 1 GB storage). Delete the stack when not in use to avoid charges.\n",
"\n",
"---"
]
},
@@ -70,14 +121,21 @@
"import os\n",
"\n",
"# Neptune cluster configuration - REPLACE WITH YOUR VALUES\n",
"# (Get these from CloudFormation stack outputs)\n",
"os.environ[\"NEPTUNE_ENDPOINT\"] = \"your-cluster.us-east-1.neptune.amazonaws.com\"\n",
"os.environ[\"NEPTUNE_PORT\"] = \"8182\"\n",
"os.environ[\"AWS_REGION\"] = \"us-east-1\"\n",
"\n",
"# AWS credentials (if using IAM Auth and not relying on IAM role or ~/.aws/credentials)\n",
"# os.environ[\"AWS_ACCESS_KEY_ID\"] = \"your-access-key-id\"\n",
"# os.environ[\"AWS_SECRET_ACCESS_KEY\"] = \"your-secret-access-key\"\n",
"# os.environ[\"AWS_SESSION_TOKEN\"] = \"your-session-token\"\n",
"# AWS credentials for IAM Authentication\n",
"# Option 1: IAM User (static credentials from CloudFormation template)\n",
"# os.environ[\"AWS_ACCESS_KEY_ID\"] = \"AKIA...\" # From AwsAccessKeyId output\n",
"# os.environ[\"AWS_SECRET_ACCESS_KEY\"] = \"...\" # From AwsSecretAccessKey output\n",
"# Note: No AWS_SESSION_TOKEN needed for IAM users\n",
"\n",
"# Option 2: IAM Role / Temporary credentials (e.g., STS AssumeRole, EC2 instance role)\n",
"# os.environ[\"AWS_ACCESS_KEY_ID\"] = \"ASIA...\" # Temporary access key\n",
"# os.environ[\"AWS_SECRET_ACCESS_KEY\"] = \"...\" # Temporary secret key\n",
"# os.environ[\"AWS_SESSION_TOKEN\"] = \"...\" # REQUIRED for temporary credentials\n",
"\n",
"print(f\"Neptune Endpoint: {os.environ.get('NEPTUNE_ENDPOINT')}\")\n",
"print(f\"AWS Region: {os.environ.get('AWS_REGION')}\")"
+228
View File
@@ -0,0 +1,228 @@
AWSTemplateFormatVersion: '2010-09-09'
Description: >
Amazon Neptune cluster with public endpoint, IAM authentication, and least-privilege
IAM user for Semantica cookbook. Uses db.t3.medium (most cost-effective Neptune instance type).
Parameters:
EnvironmentName:
Type: String
Default: semantica-neptune
Description: Environment name prefix for resource naming
Resources:
# =============================================================================
# VPC & NETWORKING
# =============================================================================
VPC:
Type: AWS::EC2::VPC
Properties:
CidrBlock: 10.0.0.0/16
EnableDnsHostnames: true
EnableDnsSupport: true
Tags:
- Key: Name
Value: !Sub ${EnvironmentName}-vpc
InternetGateway:
Type: AWS::EC2::InternetGateway
Properties:
Tags:
- Key: Name
Value: !Sub ${EnvironmentName}-igw
InternetGatewayAttachment:
Type: AWS::EC2::VPCGatewayAttachment
Properties:
InternetGatewayId: !Ref InternetGateway
VpcId: !Ref VPC
PublicSubnet1:
Type: AWS::EC2::Subnet
Properties:
VpcId: !Ref VPC
AvailabilityZone: !Select [0, !GetAZs '']
CidrBlock: 10.0.1.0/24
MapPublicIpOnLaunch: true
Tags:
- Key: Name
Value: !Sub ${EnvironmentName}-public-subnet-1
PublicSubnet2:
Type: AWS::EC2::Subnet
Properties:
VpcId: !Ref VPC
AvailabilityZone: !Select [1, !GetAZs '']
CidrBlock: 10.0.2.0/24
MapPublicIpOnLaunch: true
Tags:
- Key: Name
Value: !Sub ${EnvironmentName}-public-subnet-2
PublicRouteTable:
Type: AWS::EC2::RouteTable
Properties:
VpcId: !Ref VPC
Tags:
- Key: Name
Value: !Sub ${EnvironmentName}-public-rt
DefaultPublicRoute:
Type: AWS::EC2::Route
DependsOn: InternetGatewayAttachment
Properties:
RouteTableId: !Ref PublicRouteTable
DestinationCidrBlock: 0.0.0.0/0
GatewayId: !Ref InternetGateway
PublicSubnet1RouteTableAssociation:
Type: AWS::EC2::SubnetRouteTableAssociation
Properties:
RouteTableId: !Ref PublicRouteTable
SubnetId: !Ref PublicSubnet1
PublicSubnet2RouteTableAssociation:
Type: AWS::EC2::SubnetRouteTableAssociation
Properties:
RouteTableId: !Ref PublicRouteTable
SubnetId: !Ref PublicSubnet2
# =============================================================================
# SECURITY GROUP
# =============================================================================
NeptuneSecurityGroup:
Type: AWS::EC2::SecurityGroup
Properties:
GroupName: !Sub ${EnvironmentName}-neptune-sg
GroupDescription: Security group for Neptune cluster - allows Bolt protocol access
VpcId: !Ref VPC
SecurityGroupIngress:
- IpProtocol: tcp
FromPort: 8182
ToPort: 8182
CidrIp: 0.0.0.0/0
Description: Allow Bolt protocol access from anywhere
SecurityGroupEgress:
- IpProtocol: -1
CidrIp: 0.0.0.0/0
Description: Allow all outbound traffic
Tags:
- Key: Name
Value: !Sub ${EnvironmentName}-neptune-sg
# =============================================================================
# NEPTUNE CLUSTER
# =============================================================================
NeptuneSubnetGroup:
Type: AWS::Neptune::DBSubnetGroup
Properties:
DBSubnetGroupDescription: Subnet group for Neptune cluster
DBSubnetGroupName: !Sub ${EnvironmentName}-subnet-group
SubnetIds:
- !Ref PublicSubnet1
- !Ref PublicSubnet2
Tags:
- Key: Name
Value: !Sub ${EnvironmentName}-subnet-group
NeptuneCluster:
Type: AWS::Neptune::DBCluster
Properties:
DBClusterIdentifier: !Sub ${EnvironmentName}-cluster
DBSubnetGroupName: !Ref NeptuneSubnetGroup
VpcSecurityGroupIds:
- !Ref NeptuneSecurityGroup
EngineVersion: '1.4.6.3'
IamAuthEnabled: true
StorageEncrypted: true
DeletionProtection: false
Tags:
- Key: Name
Value: !Sub ${EnvironmentName}-cluster
NeptuneInstance:
Type: AWS::Neptune::DBInstance
Properties:
DBInstanceIdentifier: !Sub ${EnvironmentName}-instance
DBInstanceClass: db.t3.medium
DBClusterIdentifier: !Ref NeptuneCluster
PubliclyAccessible: true
Tags:
- Key: Name
Value: !Sub ${EnvironmentName}-instance
# =============================================================================
# IAM USER WITH LEAST PRIVILEGES
# =============================================================================
NeptuneUser:
Type: AWS::IAM::User
Properties:
UserName: !Sub ${EnvironmentName}-user
Tags:
- Key: Name
Value: !Sub ${EnvironmentName}-user
NeptuneUserPolicy:
Type: AWS::IAM::Policy
Properties:
PolicyName: !Sub ${EnvironmentName}-neptune-access
Users:
- !Ref NeptuneUser
PolicyDocument:
Version: '2012-10-17'
Statement:
- Sid: NeptuneDataAccess
Effect: Allow
Action:
- neptune-db:connect
- neptune-db:ReadDataViaQuery
- neptune-db:WriteDataViaQuery
- neptune-db:DeleteDataViaQuery
Resource: !Sub
- arn:aws:neptune-db:${AWS::Region}:${AWS::AccountId}:${ClusterResourceId}/*
- ClusterResourceId: !GetAtt NeptuneCluster.ClusterResourceId
NeptuneUserAccessKey:
Type: AWS::IAM::AccessKey
Properties:
UserName: !Ref NeptuneUser
# =============================================================================
# OUTPUTS
# =============================================================================
Outputs:
NeptuneEndpoint:
Description: Neptune cluster endpoint (hostname only) - use as NEPTUNE_ENDPOINT
Value: !GetAtt NeptuneCluster.Endpoint
NeptunePort:
Description: Neptune cluster port - use as NEPTUNE_PORT
Value: !GetAtt NeptuneCluster.Port
AwsAccessKeyId:
Description: Access key ID for the Neptune IAM user - use as AWS_ACCESS_KEY_ID
Value: !Ref NeptuneUserAccessKey
AwsSecretAccessKey:
Description: Secret access key for the Neptune IAM user - use as AWS_SECRET_ACCESS_KEY
Value: !GetAtt NeptuneUserAccessKey.SecretAccessKey
AwsRegion:
Description: AWS region where Neptune is deployed - use as AWS_REGION
Value: !Ref AWS::Region
NeptuneClusterResourceId:
Description: Neptune cluster resource ID (for IAM policy reference)
Value: !GetAtt NeptuneCluster.ClusterResourceId
VpcId:
Description: VPC ID
Value: !Ref VPC
SecurityGroupId:
Description: Neptune security group ID
Value: !Ref NeptuneSecurityGroup
@@ -110,7 +110,7 @@
"source": [
"# Set up API keys\n",
"# Note: In production, use environment variables: export GROQ_API_KEY=\"your-key\"\n",
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"Your Groq API\")\n"
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"\")\n"
]
},
{
@@ -30,7 +30,7 @@
"# Environment Setup\n",
"import os\n",
"\n",
"os.environ['GROQ_API_KEY'] = os.getenv('GROQ_API_KEY', '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')\n",
"os.environ['GROQ_API_KEY'] = os.getenv('GROQ_API_KEY', '')\n",
"\n",
"# Install Semantica and all required dependencies\n",
"%pip install -qU semantica networkx matplotlib plotly pandas faiss-cpu beautifulsoup4 groq sentence-transformers\n"
@@ -84,7 +84,7 @@
"source": [
"# Set up API keys\n",
"# Note: In production, use environment variables: export GROQ_API_KEY=\"your-key\"\n",
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"your-groq-api-key-here\")\n",
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"\")\n",
"\n",
"print(\"API keys configured.\")\n"
]
@@ -109,7 +109,7 @@
"source": [
"import os\n",
"\n",
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"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\")\n"
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"\")\n"
]
},
{
@@ -85,7 +85,7 @@
"source": [
"import os\n",
"\n",
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"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\")\n"
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"\")\n"
]
},
{
@@ -81,7 +81,7 @@
"source": [
"import os\n",
"\n",
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"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\")\n",
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"\")\n",
"\n",
"# Configuration constants\n",
"EMBEDDING_DIMENSION = 384\n",
@@ -98,7 +98,7 @@
"source": [
"import os\n",
"\n",
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"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\")\n",
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"\")\n",
"\n",
"# Configuration constants\n",
"EMBEDDING_DIMENSION = 384\n",
File diff suppressed because it is too large Load Diff
@@ -83,7 +83,7 @@
"source": [
"import os\n",
"\n",
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"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\")\n",
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"\")\n",
"\n",
"# Configuration constants\n",
"EMBEDDING_DIMENSION = 384\n",
@@ -80,7 +80,7 @@
"source": [
"import os\n",
"\n",
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"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\")\n",
"os.environ[\"GROQ_API_KEY\"] = os.getenv(\"GROQ_API_KEY\", \"\")\n",
"\n",
"# Configuration constants\n",
"EMBEDDING_DIMENSION = 384\n",
+5 -5
View File
@@ -12,22 +12,22 @@ How to cite Semantica in academic papers and research.
author = {Hawksight AI},
year = {2026},
url = {https://github.com/Hawksight-AI/semantica},
version = {0.2.0},
version = {0.2.3},
doi = {10.5281/zenodo.XXXXXXX}
}
```
### APA
Hawksight AI. (2026). *Semantica: An Open Source Framework for Semantic Layers and Knowledge Engineering* (Version 0.2.0) [Computer software]. https://github.com/Hawksight-AI/semantica
Hawksight AI. (2026). *Semantica: An Open Source Framework for Semantic Layers and Knowledge Engineering* (Version 0.2.3) [Computer software]. https://github.com/Hawksight-AI/semantica
### MLA
Hawksight AI. *Semantica: An Open Source Framework for Semantic Layers and Knowledge Engineering*. Version 0.2.0, GitHub, 2026, https://github.com/Hawksight-AI/semantica.
Hawksight AI. *Semantica: An Open Source Framework for Semantic Layers and Knowledge Engineering*. Version 0.2.3, GitHub, 2026, https://github.com/Hawksight-AI/semantica.
### Chicago
Hawksight AI. *Semantica: An Open Source Framework for Semantic Layers and Knowledge Engineering*. Version 0.2.0. GitHub, 2026. https://github.com/Hawksight-AI/semantica.
Hawksight AI. *Semantica: An Open Source Framework for Semantic Layers and Knowledge Engineering*. Version 0.2.3. GitHub, 2026. https://github.com/Hawksight-AI/semantica.
### IEEE
Hawksight AI, "Semantica: An Open Source Framework for Semantic Layers and Knowledge Engineering," Version 0.2.0, GitHub, 2026. [Online]. Available: https://github.com/Hawksight-AI/semantica
Hawksight AI, "Semantica: An Open Source Framework for Semantic Layers and Knowledge Engineering," Version 0.2.3, GitHub, 2026. [Online]. Available: https://github.com/Hawksight-AI/semantica
---
+1 -1
View File
@@ -17,7 +17,7 @@
<p><em>The missing fabric between raw data and AI engineering. A comprehensive open-source framework for building semantic layers and knowledge engineering systems that transform unstructured data into AI-ready knowledge — powering Knowledge Graph-Powered RAG (GraphRAG), AI Agents, Multi-Agent Systems, and AI applications with structured semantic knowledge.</em></p>
<p>🆓 <strong>100% Open Source</strong> • 📜 <strong>MIT Licensed</strong> • 🚀 <strong>Latest Version: 0.1.1</strong> • 🚀 <strong>Production Ready</strong> • 🌍 <strong>Community Driven</strong></p>
<p>🆓 <strong>100% Open Source</strong> • 📜 <strong>MIT Licensed</strong> • 🚀 <strong>Latest Version: 0.2.3</strong> • 🚀 <strong>Production Ready</strong> • 🌍 <strong>Community Driven</strong></p>
<p>
<a href="getting-started/" class="md-button md-button--primary">Get Started</a>
+6 -1
View File
@@ -9,7 +9,12 @@ document.addEventListener("DOMContentLoaded", function () {
// Define versions
var versions = [
{ name: "0.1.1", url: "#", current: true },
{ name: "0.2.4", url: "#", current: true },
{ name: "0.2.3", url: "#", current: false },
{ name: "0.2.2", url: "#", current: false },
{ name: "0.2.1", url: "#", current: false },
{ name: "0.2.0", url: "#", current: false },
{ name: "0.1.1", url: "#", current: false },
{ name: "0.1.0", url: "#", current: false }
];
+34
View File
@@ -63,6 +63,29 @@ The module uses several inference algorithms:
---
## Ontology Ingestion
Ingest existing ontology files directly into usable data structures using `OntologyIngestor`.
**Function:** `ingest_ontology(source, method="file")`
| Argument | Description |
|----------|-------------|
| `source` | File path, directory path, or list of paths |
| `method` | Ingestion method (default: "file") |
**Example:**
```python
from semantica.ontology import ingest_ontology
# Ingest file
data = ingest_ontology("ontology.ttl")
# Ingest directory
dataset = ingest_ontology("ontologies/")
```
## Main Classes
### OntologyEngine
@@ -170,6 +193,17 @@ Manages external dependencies.
| `import_external_ontology(uri, ontology)` | Load and merge external ontology |
| `evaluate_alignment(uri, ontology)` | Assess alignment and compatibility |
### OntologyIngestor
Handles ingestion of existing ontologies from files and directories.
**Methods:**
| Method | Description |
|--------|-------------|
| `ingest_ontology(file_path)` | Ingest a single ontology file |
| `ingest_directory(directory_path)` | Recursively ingest ontology files from a directory |
---
## Unified Engine Examples
+14 -6
View File
@@ -23,7 +23,7 @@ The **Semantic Extract Module** extracts structured information from unstructure
- **High Accuracy**: LLM-based extraction for complex schemas
- **Flexible Configuration**: Customize extraction for your domain
- **Confidence Scores**: Get confidence scores for all extractions
- **Batch Processing**: Efficient batch processing for large datasets
- **Batch Processing**: Efficient parallel batch processing for large datasets
- **Coreference Resolution**: Resolve pronouns to their entity references
### How It Works
@@ -185,7 +185,9 @@ Core entity extraction implementation used by notebooks and lower-level integrat
|-----------|------|---------|-------------|
| `method` | str or list | `"ml"` | Method(s): "ml", "llm", "pattern", "regex", "huggingface" |
| `silent_fail` | bool | `False` | Return empty list on error instead of raising (LLM only) |
| `max_text_length` | int | `None` | Max text length for auto-chunking (LLM only) |
| `max_text_length` | int | `64000` | Max text length for auto-chunking (LLM only) |
| `max_tokens` | int | `None` | Max output tokens for LLM generation |
| `max_workers` | int | `1` | Threads for parallel batch processing |
| `**config` | dict | `{}` | Method-specific config (e.g., `model`, `provider`) |
**Methods:**
@@ -204,11 +206,12 @@ from semantica.semantic_extract import NERExtractor
extractor = NERExtractor(method="ml", model="en_core_web_trf")
entities = extractor.extract("Elon Musk leads SpaceX.")
# 2. LLM (OpenAI/Gemini/etc)
# 2. LLM (OpenAI/Gemini/Groq/etc)
extractor = NERExtractor(
method="llm",
provider="openai",
model="gpt-4",
provider="groq",
model="llama-3.3-70b-versatile",
max_tokens=2048, # Increased output limit
temperature=0.0
)
@@ -232,6 +235,7 @@ Extracts relationships between entities.
| `bidirectional` | bool | `False` | Extract bidirectional relations |
| `confidence_threshold` | float | `0.6` | Minimum confidence score |
| `max_distance` | int | `50` | Max token distance between entities |
| `max_workers` | int | `1` | Threads for parallel batch processing |
**Methods:**
@@ -308,6 +312,7 @@ Identifies events with temporal information and participants.
| `extract_participants` | bool | `True` | Extract event participants |
| `extract_location` | bool | `True` | Extract event locations |
| `extract_time` | bool | `True` | Extract temporal information |
| `max_workers` | int | `1` | Threads for parallel batch processing |
**Methods:**
@@ -340,7 +345,9 @@ Extracts RDF triplets (Subject-Predicate-Object).
| `include_provenance` | bool | `False` | Track source sentences |
| `method` | str | `"pattern"` | Extraction method ("pattern", "rules", "huggingface", "llm") |
| `silent_fail` | bool | `False` | Return empty list on error instead of raising (LLM only) |
| `max_text_length` | int | `None` | Max text length for auto-chunking (LLM only) |
| `max_text_length` | int | `64000` | Max text length for auto-chunking (LLM only) |
| `max_tokens` | int | `None` | Max output tokens for LLM generation |
| `max_workers` | int | `1` | Threads for parallel batch processing |
**Methods:**
@@ -371,6 +378,7 @@ Extracts structured semantic networks with nodes and edges.
|-----------|------|---------|-------------|
| `ner_method` | str | `None` | Method for node extraction |
| `relation_method` | str | `None` | Method for edge extraction |
| `max_workers` | int | `1` | Threads for parallel batch processing |
| `**config` | dict | `{}` | Configuration for underlying extractors |
**Methods:**
+14 -2
View File
@@ -112,6 +112,8 @@ The main facade for all vector operations.
| Method | Description |
|--------|-------------|
| `store_vectors(vectors, metadata)` | Store embeddings |
| `add_documents(documents, metadata, batch_size, parallel)` | **(New)** Store documents with automatic embedding generation and parallelization |
| `embed_batch(texts)` | **(New)** Generate embeddings for a batch of texts |
| `search(query, k)` | Semantic search |
| `delete(ids)` | Remove vectors |
@@ -120,15 +122,24 @@ The main facade for all vector operations.
```python
from semantica.vector_store import VectorStore
# Initialize (defaults to FAISS)
# Initialize (defaults to FAISS, parallel enabled by default with 6 workers)
store = VectorStore(backend="faiss", dimension=1536)
# Store
# 1. Store pre-computed vectors
ids = store.store_vectors(
vectors=[[0.1, 0.2, ...], ...],
metadata=[{"text": "Hello"}, ...]
)
# 2. Store raw documents (High Performance)
# Automatically handles embedding generation in parallel batches (uses default 6 workers)
ids = store.add_documents(
documents=["Doc 1", "Doc 2", ...],
metadata=[{"id": 1}, {"id": 2}, ...],
batch_size=32,
parallel=True
)
# Search
results = store.search(query_vector=[0.1, 0.2, ...], k=5)
```
@@ -771,6 +782,7 @@ print(f"Context: {context}")
---
## See Also
- [High-Performance Usage Guide](../vector_store_usage.md) - **(New)** Parallel ingestion and batching guide
- [Embeddings Module](embeddings.md) - Generates the vectors
- [Context Module](context.md) - Uses vector store for memory
- [Ingest Module](ingest.md) - Source of data
+101
View File
@@ -0,0 +1,101 @@
# High-Performance Vector Store Usage
This guide demonstrates how to leverage the new high-performance features of the Semantica Vector Store, specifically designed for efficient batch processing and parallel ingestion of large document sets.
## 🚀 Key Features
- **Parallel Ingestion**: Utilize multi-threading to embed and store documents concurrently.
- **Batch Processing**: Automatically group documents into batches to minimize overhead.
- **Unified API**: A single `add_documents` method handles embedding generation and storage.
---
## ⚡ Quick Start: Parallel Ingestion
The fastest way to ingest documents is using the `add_documents` method. Parallelization is enabled by default with optimized settings (6 workers).
```python
from semantica.vector_store import VectorStore
import time
store = VectorStore(
backend="faiss",
dimension=768,
)
documents = [f"This is document number {i} with some content." for i in range(1000)]
metadata = [{"source": "generated", "id": i} for i in range(1000)]
start_time = time.time()
ids = store.add_documents(
documents=documents,
metadata=metadata,
batch_size=64,
parallel=True,
)
print(f"Ingested {len(ids)} documents in {time.time() - start_time:.2f}s")
```
---
## 📊 Performance Comparison
### Old Method (Sequential Loop)
*Slower due to sequential processing and overhead per single item.*
```python
for doc in documents:
emb = embedder.generate(doc)
store.store_vectors([emb], [{"text": doc}])
```
### New Method (Parallel Batching)
*Significantly faster (3x-10x) by utilizing thread pools and batch operations.*
```python
store.add_documents(documents, parallel=True)
```
---
## 🛠 Configuration & Tuning
### `max_workers`
Controls the number of concurrent threads used for embedding generation.
- **Default**: 6 (Optimized for most systems)
- **Recommendation**: You generally don't need to change this. If you have very high core counts or specific throughput needs, you can override it.
```python
store = VectorStore(max_workers=16)
```
### `batch_size`
Controls how many documents are processed in a single chunk.
- **Default**: 32
- **Recommendation**:
- **Local Models**: 32-64 usually works well.
- **API Models (OpenAI, etc.)**: Larger batches (e.g., 100-200) can reduce network latency overhead.
```python
store.add_documents(documents, batch_size=100)
```
---
## 🧩 Advanced: Manual Batch Embedding
If you need the embeddings without storing them immediately, use `embed_batch`.
```python
vectors = store.embed_batch(
texts=documents[:100],
)
print(f"Generated {len(vectors)} vectors")
```
## ⚠️ Best Practices
1. **Metadata Consistency**: Ensure your `metadata` list has the same length as your `documents` list.
2. **Error Handling**: The `add_documents` method will propagate exceptions if embedding fails. Ensure your data is clean.
3. **Memory Usage**: Very large `batch_size` combined with high `max_workers` can increase memory usage. Monitor your system resources.
+165 -284
View File
@@ -4,317 +4,198 @@ build-backend = "setuptools.build_meta"
[project]
name = "semantica"
version = "0.2.0"
description = "🧠 Semantica - An Open Source Framework for building Semantic Layers and Knowledge Engineering "
version = "0.2.4"
description = "🧠 Semantica - An Open Source Framework for building Semantic Layers and Knowledge Engineering"
readme = "README.md"
license = {text = "MIT"}
authors = [
{name = "Hawksight AI", email = "semantica-dev@users.noreply.github.com"}
]
maintainers = [
{name = "Hawksight AI", email = "semantica-dev@users.noreply.github.com"}
]
license = { text = "MIT" }
authors = [{ name = "Hawksight AI", email = "semantica-dev@users.noreply.github.com" }]
maintainers = [{ name = "Hawksight AI", email = "semantica-dev@users.noreply.github.com" }]
requires-python = ">=3.8"
classifiers = [
"Development Status :: 3 - Alpha",
"Intended Audience :: Developers",
"Intended Audience :: Science/Research",
"License :: OSI Approved :: MIT License",
"Operating System :: OS Independent",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.8",
"Programming Language :: Python :: 3.9",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Topic :: Scientific/Engineering :: Artificial Intelligence",
"Topic :: Software Development :: Libraries :: Python Modules",
"Topic :: Text Processing :: Linguistic",
"Topic :: Database :: Database Engines/Servers",
"Topic :: Internet :: WWW/HTTP :: Indexing/Search"
"Development Status :: 3 - Alpha",
"Intended Audience :: Developers",
"Intended Audience :: Science/Research",
"License :: OSI Approved :: MIT License",
"Operating System :: OS Independent",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.8",
"Programming Language :: Python :: 3.9",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Topic :: Scientific/Engineering :: Artificial Intelligence",
"Topic :: Software Development :: Libraries :: Python Modules"
]
keywords = [
"semantic-layer", "knowledge-engineering", "nlp", "knowledge-graph",
"embeddings", "entity-extraction", "relationship-extraction", "rdf",
"ontology", "semantic-analysis", "ai", "machine-learning"
"semantic-layer", "knowledge-graph", "nlp", "embeddings",
"entity-extraction", "relationship-extraction", "rdf", "ontology"
]
# ---------------- CORE DEPENDENCIES (SAFE DEFAULT) ----------------
dependencies = [
"numpy>=1.21.0",
"pandas>=1.3.0",
"scikit-learn>=1.0.0",
"umap-learn>=0.5.0",
"spacy>=3.4.0",
"transformers>=4.20.0",
"torch>=1.12.0",
"sentence-transformers>=2.2.0",
"rdflib>=6.2.0",
"networkx>=2.8.0",
"matplotlib>=3.5.0",
"seaborn>=0.11.0",
"plotly>=5.10.0",
"ipywidgets>=8.0.0",
"requests>=2.28.0",
"GitPython>=3.1.30",
"chardet>=5.1.0",
"protobuf==4.25.8",
"grpcio==1.67.1",
"beautifulsoup4>=4.11.0",
"lxml>=4.9.0",
"pypdf2>=2.10.0",
"python-docx>=0.8.11",
"docling>=1.0.0",
"openpyxl>=3.0.10",
"pillow>=9.2.0",
"librosa>=0.9.0",
"opencv-python>=4.6.0",
"faiss-cpu>=1.7.0",
"fastembed>=0.2.0",
"onnxruntime>=1.17.0",
"tokenizers>=0.15.0",
"weaviate-client>=3.15.0",
"qdrant-client>=1.3.0",
"neo4j>=5.0.0",
"falkordb>=1.0.0",
"pymongo>=4.2.0",
"sqlalchemy>=1.4.0",
"psycopg2-binary>=2.9.0",
"pymysql>=1.0.0",
"redis>=4.3.0",
"celery>=5.2.0",
"kafka-python>=2.0.0",
"pulsar-client>=3.0.0",
"pika>=1.3.0",
"boto3>=1.24.0",
"azure-storage-blob>=12.12.0",
"google-cloud-storage>=2.5.0",
"pydantic>=2.0.0",
"fastmcp>=0.1.0",
"groq>=0.4.0",
"openai>=1.0.0",
"litellm>=1.0.0",
"instructor>=1.0.0",
"click>=8.1.0",
"rich>=12.5.0",
"tqdm>=4.64.0",
"pyyaml>=6.0",
"toml>=0.10.0",
"python-dotenv>=0.20.0",
"loguru>=0.6.0",
"structlog>=22.1.0",
"prometheus-client>=0.14.0",
"opentelemetry-api>=1.12.0",
"opentelemetry-sdk>=1.12.0",
"opentelemetry-instrumentation",
"fastapi>=0.78.0",
"uvicorn>=0.18.0",
"pytest>=7.1.0",
"pytest-cov>=3.0.0",
"pytest-asyncio>=0.19.0",
"black>=22.6.0",
"isort>=5.10.0",
"flake8>=4.0.0",
"mypy>=0.971",
"pre-commit>=2.19.0"
"numpy>=1.21.0",
"pandas>=1.3.0",
"scikit-learn>=1.0.0",
"umap-learn>=0.5.0",
"spacy>=3.4.0",
"transformers>=4.20.0",
"torch>=1.12.0",
"sentence-transformers>=2.2.0",
"rdflib>=6.2.0",
"networkx>=2.8.0",
"matplotlib>=3.5.0",
"seaborn>=0.11.0",
"plotly>=5.10.0",
"ipywidgets>=8.0.0",
"requests>=2.28.0",
"GitPython>=3.1.30",
"chardet>=5.1.0",
"protobuf>=5.29.1,<7.0",
"grpcio>=1.71.2",
"beautifulsoup4>=4.11.0",
"lxml>=4.9.0",
"pypdf2>=2.10.0",
"python-docx>=0.8.11",
"openpyxl>=3.0.10",
"pillow>=9.2.0",
"librosa>=0.9.0",
"opencv-python>=4.6.0",
"faiss-cpu>=1.7.0",
"fastembed>=0.2.0",
"onnxruntime>=1.17.0",
"tokenizers>=0.15.0",
"pydantic>=2.0.0",
"click>=8.1.0",
"rich>=12.5.0",
"tqdm>=4.64.0",
"pyyaml>=6.0",
"toml>=0.10.0",
"python-dotenv>=0.20.0",
"loguru>=0.6.0",
"structlog>=22.1.0"
]
[project.urls]
Homepage = "https://github.com/Hawksight-AI/semantica"
Repository = "https://github.com/Hawksight-AI/semantica"
"Bug Tracker" = "https://github.com/Hawksight-AI/semantica/issues"
Discussions = "https://github.com/Hawksight-AI/semantica/discussions"
# ---------------- OPTIONAL DEPENDENCIES ----------------
[project.optional-dependencies]
dev = [
"pytest>=7.1.0",
"pytest-cov>=3.0.0",
"pytest-asyncio>=0.19.0",
"black>=22.6.0",
"isort>=5.10.0",
"flake8>=4.0.0",
"mypy>=0.971",
"pre-commit>=2.19.0",
"jupyter>=1.0.0",
"ipykernel>=6.15.0",
"notebook>=6.4.0"
]
viz = [
"pyvis>=0.3.0",
"graphviz>=0.20.0",
"umap-learn>=0.5.0",
"d3blocks>=1.0.0"
]
gpu = [
"torch>=1.12.0",
"faiss-gpu>=1.7.0",
"cupy>=10.0.0"
]
cloud = [
"boto3>=1.24.0",
"azure-storage-blob>=12.12.0",
"google-cloud-storage>=2.5.0",
"kubernetes>=24.0.0",
"helm>=3.10.0"
]
monitoring = [
"prometheus-client>=0.14.0",
"opentelemetry-api>=1.12.0",
"opentelemetry-sdk>=1.12.0",
"opentelemetry-instrumentation>=0.32.0",
"grafana-api>=1.0.0",
"elasticsearch>=8.5.0"
]
llm-openai = [
"openai>=1.0.0"
]
llm-gemini = [
"google-generativeai>=0.3.0"
]
llm-groq = [
"groq>=0.4.0"
]
llm-anthropic = [
"anthropic>=0.18.0"
]
llm-ollama = [
"ollama>=0.1.0"
]
llm-deepseek = [
"deepseek>=0.1.0"
]
llm-litellm = [
"litellm>=1.0.0"
]
llm-instructor = [
"instructor>=1.0.0"
]
# ---- LLM Providers ----
llm-openai = ["openai>=1.0.0"]
llm-groq = ["groq>=0.4.0"]
llm-gemini = ["google-genai>=0.1.0"]
llm-anthropic = ["anthropic>=0.18.0"]
llm-ollama = ["ollama>=0.1.0"]
llm-deepseek = ["deepseek>=0.1.0"]
llm-litellm = ["litellm>=1.0.0"]
llm-instructor = ["instructor>=1.0.0"]
llm-all = [
"semantica[llm-openai,llm-gemini,llm-groq,llm-anthropic,llm-ollama,llm-deepseek,llm-litellm,llm-instructor]"
]
models-huggingface = [
"transformers>=4.20.0",
"torch>=1.12.0"
]
split-tiktoken = [
"tiktoken>=0.5.0"
]
split-community = [
"python-louvain>=0.16"
]
split-topic = [
"bertopic>=0.15.0",
"gensim>=4.3.0"
]
split-all = [
"semantica[split-tiktoken,split-community,split-topic]"
]
graph-neo4j = [
"neo4j>=5.0.0"
]
graph-falkordb = [
"falkordb>=1.0.0",
"redis>=4.3.0"
]
graph-amazon-neptune = [
"boto3>=1.24.0",
"neo4j>=5.0.0"
]
graph-all = [
"semantica[graph-neo4j,graph-falkordb,graph-amazon-neptune]"
]
parse-docling = [
"docling>=1.0.0"
]
all = [
"semantica[dev,viz,gpu,cloud,monitoring,llm-all,models-huggingface,split-all,graph-all,parse-docling]"
"semantica[llm-openai,llm-groq,llm-gemini,llm-anthropic,llm-ollama,llm-deepseek,llm-litellm,llm-instructor]"
]
# ---- Document Parsing ----
parse-docling = ["docling>=1.0.0"]
# ---- Embedding / Models ----
models-huggingface = [
"transformers>=4.20.0",
"torch>=1.12.0"
]
# ---- Graph Backends ----
graph-neo4j = ["neo4j>=5.0.0"]
graph-falkordb = ["falkordb>=1.0.0", "redis>=4.3.0"]
graph-amazon-neptune = ["boto3>=1.24.0", "neo4j>=5.0.0"]
graph-all = [
"semantica[graph-neo4j,graph-falkordb,graph-amazon-neptune]"
]
# ---- Infra / Queues / Workers ----
infra = [
"redis>=4.3.0",
"celery>=5.2.0",
"kafka-python>=2.0.0",
"pulsar-client>=3.0.0",
"pika>=1.3.0"
]
# ---- Cloud Providers ----
cloud = [
"boto3>=1.24.0",
"azure-storage-blob>=12.12.0",
"google-cloud-storage>=2.5.0"
]
# ---- Monitoring (FIXED) ----
monitoring = [
"prometheus-client>=0.14.0",
"opentelemetry-api>=1.30.0,<2.0.0",
"opentelemetry-sdk>=1.30.0,<2.0.0",
"opentelemetry-semantic-conventions>=0.58b0,<0.61b0",
"opentelemetry-instrumentation>=0.58b0,<0.61b0"
]
# ---- Visualization ----
viz = [
"pyvis>=0.3.0",
"graphviz>=0.20.0",
"d3blocks>=1.0.0"
]
# ---- GPU ----
gpu = [
"faiss-gpu>=1.7.0",
"cupy>=10.0.0"
]
# ---- Splitting / Chunking ----
split-tiktoken = ["tiktoken>=0.5.0"]
split-community = ["python-louvain>=0.16"]
split-topic = ["bertopic>=0.15.0", "gensim>=4.3.0"]
split-all = [
"semantica[split-tiktoken,split-community,split-topic]"
]
# ---- Dev ----
dev = [
"pytest>=7.1.0",
"pytest-cov>=3.0.0",
"pytest-asyncio>=0.19.0",
"black>=22.6.0",
"isort>=5.10.0",
"flake8>=4.0.0",
"mypy>=0.971",
"pre-commit>=2.19.0",
"jupyter>=1.0.0",
"ipykernel>=6.15.0"
]
# ---- Everything ----
all = [
"semantica[dev,viz,gpu,infra,cloud,monitoring,llm-all,models-huggingface,split-all,graph-all,parse-docling]"
]
# ---------------- ENTRYPOINTS ----------------
[project.scripts]
semantica = "semantica.cli:main"
semantica-server = "semantica.server:main"
semantica-worker = "semantica.worker:main"
# ---------------- TOOLING ----------------
[tool.setuptools.packages.find]
where = ["."]
include = ["semantica*"]
exclude = ["tests*", "docs*", "examples*"]
[tool.setuptools.package-data]
semantica = ["*.yaml", "*.yml", "*.json", "*.toml", "*.txt", "*.md"]
[tool.black]
line-length = 88
target-version = ['py38', 'py39', 'py310', 'py311', 'py312']
include = '\.pyi?$'
extend-exclude = '''
/(
# directories
\.eggs
| \.git
| \.hg
| \.mypy_cache
| \.tox
| \.venv
| build
| dist
)/
'''
[tool.isort]
profile = "black"
multi_line_output = 3
line_length = 88
known_first_party = ["semantica"]
known_third_party = ["numpy", "pandas", "scikit-learn", "spacy", "transformers", "torch"]
[tool.mypy]
python_version = "3.9"
warn_return_any = true
warn_unused_configs = true
disallow_untyped_defs = true
disallow_incomplete_defs = true
check_untyped_defs = true
disallow_untyped_decorators = true
no_implicit_optional = true
warn_redundant_casts = true
warn_unused_ignores = true
warn_no_return = true
warn_unreachable = true
strict_equality = true
show_error_codes = true
[tool.pytest.ini_options]
minversion = "7.0"
addopts = "-ra -q --strict-markers --strict-config"
testpaths = ["tests"]
python_files = ["test_*.py", "*_test.py"]
python_classes = ["Test*"]
python_functions = ["test_*"]
markers = [
"slow: marks tests as slow (deselect with '-m \"not slow\"')",
"integration: marks tests as integration tests",
"unit: marks tests as unit tests",
"gpu: marks tests that require GPU",
"cloud: marks tests that require cloud services"
]
[tool.coverage.run]
source = ["semantica"]
omit = [
"*/tests/*",
"*/test_*",
"*/__pycache__/*",
"*/migrations/*"
]
[tool.coverage.report]
exclude_lines = [
"pragma: no cover",
"def __repr__",
"if self.debug:",
"if settings.DEBUG",
"raise AssertionError",
"raise NotImplementedError",
"if 0:",
"if __name__ == .__main__.:",
"class .*\\bProtocol\\):",
"@(abc\\.)?abstractmethod"
]
+1 -1
View File
@@ -10,7 +10,7 @@ Main exports:
- Config: Configuration management
"""
__version__ = "0.2.0"
__version__ = "0.2.4"
__author__ = "Semantica Contributors"
__license__ = "MIT"
+9
View File
@@ -86,6 +86,7 @@ Main Classes:
- RepoIngestor: Git repository processing
- EmailIngestor: Email protocol handling
- DBIngestor: Database export handling
- OntologyIngestor: Ontology file processing
- MethodRegistry: Registry for custom ingestion methods
- IngestConfig: Configuration manager for ingest module
@@ -98,6 +99,7 @@ Convenience Functions:
- ingest_repository: Repository ingestion wrapper
- ingest_email: Email ingestion wrapper
- ingest_database: Database ingestion wrapper
- ingest_ontology: Ontology ingestion wrapper
Example Usage:
@@ -134,6 +136,7 @@ from .methods import (
ingest_feed,
ingest_file,
ingest_mcp,
ingest_ontology,
ingest_repository,
ingest_stream,
ingest_web,
@@ -166,6 +169,8 @@ from .web_ingestor import (
WebIngestor,
)
from .ontology_ingestor import OntologyData, OntologyIngestor
__all__ = [
# File ingestion
"FileIngestor",
@@ -216,6 +221,9 @@ __all__ = [
"MCPClient",
"MCPResource",
"MCPTool",
# Ontology ingestion
"OntologyIngestor",
"OntologyData",
# Registry and Methods
"MethodRegistry",
"method_registry",
@@ -227,6 +235,7 @@ __all__ = [
"ingest_repository",
"ingest_email",
"ingest_database",
"ingest_ontology",
"ingest_mcp",
"get_ingest_method",
"list_available_methods",
+69
View File
@@ -150,6 +150,7 @@ from .email_ingestor import EmailData, EmailIngestor
from .feed_ingestor import FeedData, FeedIngestor
from .file_ingestor import FileIngestor, FileObject
from .mcp_ingestor import MCPData, MCPIngestor
from .ontology_ingestor import OntologyData, OntologyIngestor
from .registry import method_registry
from .repo_ingestor import RepoIngestor
from .stream_ingestor import StreamIngestor, StreamProcessor
@@ -537,6 +538,66 @@ def ingest_email(
raise
def ingest_ontology(
source: Union[str, Path, List[Union[str, Path]]], method: str = "file", **kwargs
) -> Union[OntologyData, List[OntologyData]]:
"""
Ingest ontology from source (convenience function).
This is a user-friendly wrapper that ingests ontologies using the specified method.
Args:
source: Ontology file path, directory path, or list of paths
method: Ingestion method (default: "file")
- "file": Single file ingestion
- "directory": Directory ingestion with recursive scanning
**kwargs: Additional options passed to OntologyIngestor
Returns:
OntologyData, List[OntologyData] with ingestion results
Examples:
>>> from semantica.ingest.methods import ingest_ontology
>>> ontology = ingest_ontology("ontology.ttl")
>>> ontologies = ingest_ontology("./ontologies", method="directory")
"""
# Check for custom method in registry
custom_method = method_registry.get("ontology", method)
if custom_method and custom_method != ingest_ontology:
try:
return custom_method(source, **kwargs)
except Exception as e:
logger.warning(
f"Custom method {method} failed: {e}, falling back to default"
)
try:
# Get config
config = ingest_config.get_method_config("ontology")
config.update(kwargs)
ingestor = OntologyIngestor(**config)
source_path = str(source) if isinstance(source, (str, Path)) else None
if method == "file" and source_path:
if isinstance(source, list):
return [ingestor.ingest_ontology(str(s), **kwargs) for s in source]
return ingestor.ingest_ontology(source_path, **kwargs)
elif method == "directory" and source_path:
recursive = kwargs.get("recursive", ingest_config.get("recursive", True))
return ingestor.ingest_directory(source_path, recursive=recursive, **kwargs)
else:
# Default: try as file
if isinstance(source, list):
return [ingestor.ingest_ontology(str(s), **kwargs) for s in source]
return ingestor.ingest_ontology(str(source), **kwargs)
except Exception as e:
logger.error(f"Failed to ingest ontology: {e}")
raise
def ingest_database(
source: Union[str, Dict[str, Any]], method: Optional[str] = None, **kwargs
) -> Union[TableData, List[TableData], Dict[str, Any]]:
@@ -769,6 +830,7 @@ def ingest(
- "repo": Repository ingestion
- "email": Email ingestion
- "db": Database ingestion
- "ontology": Ontology ingestion
method: Optional specific ingestion method
**kwargs: Additional options passed to ingestor
@@ -802,6 +864,8 @@ def ingest(
("git@", "https://github.com", "https://gitlab.com")
):
source_type = "repo"
elif source_str.endswith((".ttl", ".owl", ".rdf", ".jsonld", ".n3", ".nt")):
source_type = "ontology"
else:
source_type = "file"
else:
@@ -830,6 +894,8 @@ def ingest(
raise ProcessingError("Email ingestion requires configuration dictionary")
elif source_type == "db":
return {"data": ingest_database(sources, method=method, **kwargs)}
elif source_type == "ontology":
return {"ontology": ingest_ontology(sources, method=method or "file", **kwargs)}
elif source_type == "mcp":
return {"data": ingest_mcp(sources, method=method or "resources", **kwargs)}
else:
@@ -909,5 +975,8 @@ method_registry.register("mcp", "default", ingest_mcp)
method_registry.register("mcp", "resources", ingest_mcp)
method_registry.register("mcp", "tools", ingest_mcp)
method_registry.register("mcp", "all", ingest_mcp)
method_registry.register("ontology", "default", ingest_ontology)
method_registry.register("ontology", "file", ingest_ontology)
method_registry.register("ontology", "directory", ingest_ontology)
method_registry.register("ingest", "default", ingest)
method_registry.register("ingest", "unified", ingest)
+392
View File
@@ -0,0 +1,392 @@
"""
Ontology Ingestion Module
This module provides capabilities to ingest external ontologies from files (OWL, RDF, TTL, etc.)
and convert them into Semantica's internal ontology dictionary format.
Supported Formats:
- Turtle (.ttl): Terse RDF Triple Language. A concise, human-readable
format for representing RDF graphs. Commonly used for writing
ontologies by hand.
- RDF/XML (.rdf, .owl): The XML serialization of RDF. The standard
format for OWL (Web Ontology Language) ontologies and often used
for data interchange.
- JSON-LD (.jsonld): JSON for Linked Data. A lightweight Linked Data
format that is easy for humans to read and for machines to parse
and generate. Ideal for web-based applications.
- N-Triples (.nt): A line-based, plain text format for encoding an
RDF graph. Each line represents a single triple. Very simple to
parse but verbose.
- Notation3 (.n3): A superset of Turtle that adds features like logic
and rules.
Key Features:
- Support for multiple RDF formats (Turtle, RDF/XML, JSON-LD, N3, NT)
- Automatic parsing using rdflib
- Conversion to Semantica ontology structure
- Batch processing of ontology files
- Extraction of classes, properties, and metadata
Example Usage:
>>> from semantica.ingest import OntologyIngestor
>>> ingestor = OntologyIngestor()
>>> ontology = ingestor.ingest_ontology("my_ontology.ttl")
"""
import os
from dataclasses import dataclass, field
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, List, Optional, Union
import rdflib
from rdflib import RDF, RDFS, OWL, Graph
from ..utils.exceptions import ProcessingError, ValidationError
from ..utils.logging import get_logger
from ..utils.progress_tracker import get_progress_tracker
@dataclass
class OntologyData:
"""Ontology data representation."""
data: Dict[str, Any]
source_path: str
format: str
metadata: Dict[str, Any] = field(default_factory=dict)
ingested_at: datetime = field(default_factory=datetime.now)
class OntologyIngestor:
"""
Ontology ingestion handler.
This class parses OWL/RDF files and converts them to Semantica's ontology dictionary format.
"""
def __init__(self, config: Optional[Dict[str, Any]] = None, **kwargs):
"""
Initialize ontology ingestor.
Args:
config: Optional configuration dictionary
**kwargs: Additional configuration parameters
"""
self.logger = get_logger("ontology_ingestor")
self.progress = get_progress_tracker()
self.config = config or {}
self.config.update(kwargs)
def ingest_ontology(self, file_path: Union[str, Path], format: Optional[str] = None, **kwargs) -> OntologyData:
"""
Ingest an ontology file.
Args:
file_path: Path to the ontology file (string or Path object)
format: Optional format hint (e.g., 'turtle', 'xml'). If None, rdflib guesses.
**kwargs: Additional arguments for rdflib parsing
Returns:
OntologyData object containing the parsed ontology and metadata
"""
file_path = Path(file_path)
# Track file ingestion
tracking_id = self.progress.start_tracking(
file=str(file_path),
module="ingest",
submodule="OntologyIngestor",
message=f"Ontology: {file_path.name}",
)
try:
# Validate file exists
if not file_path.exists():
raise ValidationError(f"File not found: {file_path}")
self.progress.update_tracking(tracking_id, message="Parsing RDF graph...")
g = Graph()
# Use provided format or let rdflib guess based on extension
parse_kwargs = kwargs.copy()
if format:
parse_kwargs['format'] = format
try:
g.parse(file_path, **parse_kwargs)
except Exception as e:
# Fallback: try to guess format from extension if not provided and initial parse failed
if not format:
ext = os.path.splitext(file_path)[1].lower()
fmt_map = {
'.ttl': 'turtle',
'.owl': 'xml', # OWL is often XML
'.rdf': 'xml',
'.jsonld': 'json-ld',
'.n3': 'n3',
'.nt': 'nt'
}
guessed_fmt = fmt_map.get(ext)
if guessed_fmt:
self.logger.info(f"Retrying with guessed format: {guessed_fmt}")
g.parse(file_path, format=guessed_fmt, **kwargs)
else:
raise e
else:
raise e
self.progress.update_tracking(tracking_id, message="Converting to internal format...")
# Determine format for metadata
used_format = format
if not used_format:
ext = os.path.splitext(file_path)[1].lower()
fmt_map = {
'.ttl': 'turtle',
'.owl': 'xml',
'.rdf': 'xml',
'.jsonld': 'json-ld',
'.n3': 'n3',
'.nt': 'nt'
}
used_format = fmt_map.get(ext, 'unknown')
ontology_dict = self._convert_to_dict(g, source_path=str(file_path), format=used_format)
ontology_data = OntologyData(
data=ontology_dict,
source_path=str(file_path),
format=used_format,
metadata=ontology_dict.get("metadata", {}).copy()
)
self.progress.stop_tracking(
tracking_id,
status="completed",
message=f"Successfully ingested ontology from {file_path}",
)
return ontology_data
except Exception as e:
self.logger.error(f"Failed to ingest ontology: {str(e)}")
self.progress.stop_tracking(
tracking_id, status="failed", message=str(e)
)
raise ProcessingError(f"Failed to ingest ontology: {str(e)}") from e
def ingest_directory(self, directory_path: Union[str, Path], recursive: bool = True, **kwargs) -> List[OntologyData]:
"""
Ingest all ontology files in a directory.
Args:
directory_path: Path to the directory (string or Path object)
recursive: Whether to search recursively
**kwargs: Additional arguments
Returns:
List of OntologyData objects
"""
directory_path = Path(directory_path)
ontologies = []
extensions = {'.ttl', '.owl', '.rdf', '.jsonld', '.n3', '.nt'}
# Track directory ingestion
tracking_id = self.progress.start_tracking(
file=str(directory_path),
module="ingest",
submodule="OntologyIngestor",
message=f"Directory: {directory_path.name}",
)
try:
if not directory_path.exists():
raise ValidationError(f"Directory not found: {directory_path}")
if not directory_path.is_dir():
raise ValidationError(f"Path is not a directory: {directory_path}")
files_to_process = []
for root, _, files in os.walk(directory_path):
for file in files:
ext = os.path.splitext(file)[1].lower()
if ext in extensions:
files_to_process.append(os.path.join(root, file))
if not recursive:
break
total_files = len(files_to_process)
self.progress.update_tracking(
tracking_id, message=f"Processing {total_files} ontology files"
)
for idx, file_path in enumerate(files_to_process, 1):
try:
ont_data = self.ingest_ontology(file_path, **kwargs)
ontologies.append(ont_data)
self.progress.update_progress(
tracking_id,
processed=idx,
total=total_files,
message=f"Processing {idx}/{total_files}: {Path(file_path).name}"
)
except Exception as e:
self.logger.warning(f"Skipping {file_path}: {e}")
self.progress.stop_tracking(
tracking_id,
status="completed",
message=f"Ingested {len(ontologies)} ontologies",
)
return ontologies
except Exception as e:
self.progress.stop_tracking(
tracking_id, status="failed", message=str(e)
)
raise
def _convert_to_dict(self, graph: Graph, source_path: str, format: str = "unknown") -> Dict[str, Any]:
"""
Convert rdflib Graph to Semantica ontology dictionary.
Args:
graph: Parsed rdflib Graph
source_path: Source file path
format: Format of the ontology file
Returns:
Ontology dictionary
"""
ontology = {
"uri": "",
"name": os.path.basename(source_path),
"version": "1.0",
"classes": [],
"properties": [],
"metadata": {
"source_path": source_path,
"ingested_at": datetime.now().isoformat(),
"format": format
}
}
# 1. Extract Ontology Metadata
for s, p, o in graph.triples((None, RDF.type, OWL.Ontology)):
ontology["uri"] = str(s)
# Try to find label/comment/versionInfo
for _, _, label in graph.triples((s, RDFS.label, None)):
ontology["name"] = str(label)
for _, _, comment in graph.triples((s, RDFS.comment, None)):
ontology["description"] = str(comment)
for _, _, version in graph.triples((s, OWL.versionInfo, None)):
ontology["version"] = str(version)
# Break after first ontology definition found (usually only one per file)
break
# 2. Extract Classes
classes = {}
# Union of owl:Class and rdfs:Class
class_types = [OWL.Class, RDFS.Class]
for c_type in class_types:
for s, p, o in graph.triples((None, RDF.type, c_type)):
if isinstance(s, rdflib.BNode):
continue # Skip blank nodes for now
uri = str(s)
if uri not in classes:
cls_def = {
"uri": uri,
"name": self._get_local_name(uri),
"type": "class"
}
# Add label/comment
label = graph.value(s, RDFS.label)
if label:
cls_def["label"] = str(label)
cls_def["name"] = str(label) # Prefer label as name if available? Or keep URI fragment?
# Keeping local name from URI is safer for internal IDs, label for display.
# But Semantica seems to use "name" for the identifier in some examples.
# Let's keep name as local name or label if simple.
comment = graph.value(s, RDFS.comment)
if comment:
cls_def["description"] = str(comment)
# Superclasses
parents = []
for _, _, parent in graph.triples((s, RDFS.subClassOf, None)):
if isinstance(parent, rdflib.URIRef):
parents.append(str(parent))
if parents:
cls_def["parents"] = parents
classes[uri] = cls_def
ontology["classes"] = list(classes.values())
# 3. Extract Properties
properties = {}
# Object Properties
for s, p, o in graph.triples((None, RDF.type, OWL.ObjectProperty)):
self._add_property(graph, s, "object", properties)
# Datatype Properties
for s, p, o in graph.triples((None, RDF.type, OWL.DatatypeProperty)):
self._add_property(graph, s, "data", properties)
# RDF Properties (generic)
for s, p, o in graph.triples((None, RDF.type, RDF.Property)):
if str(s) not in properties: # Don't overwrite if already found as specific type
self._add_property(graph, s, "annotation", properties) # Default to annotation or generic
ontology["properties"] = list(properties.values())
return ontology
def _add_property(self, graph: Graph, subject: rdflib.term.Node, prop_type: str, properties_dict: Dict):
if isinstance(subject, rdflib.BNode):
return
uri = str(subject)
if uri in properties_dict:
return
prop_def = {
"uri": uri,
"name": self._get_local_name(uri),
"type": prop_type
}
label = graph.value(subject, RDFS.label)
if label:
prop_def["label"] = str(label)
comment = graph.value(subject, RDFS.comment)
if comment:
prop_def["description"] = str(comment)
# Domain and Range
domain = graph.value(subject, RDFS.domain)
if domain and isinstance(domain, rdflib.URIRef):
prop_def["domain"] = str(domain)
range_val = graph.value(subject, RDFS.range)
if range_val and isinstance(range_val, rdflib.URIRef):
prop_def["range"] = str(range_val)
properties_dict[uri] = prop_def
def _get_local_name(self, uri: str) -> str:
"""Extract local name from URI."""
if '#' in uri:
return uri.split('#')[-1]
return uri.split('/')[-1]
+51 -5
View File
@@ -161,7 +161,14 @@ class GraphBuilder:
}
all_relationships.append(rel_dict)
elif isinstance(item, dict):
# Detect and normalize Entity objects inside dict
if "source_id" in item and "source" not in item:
item["source"] = item["source_id"]
if "target_id" in item and "target" not in item:
item["target"] = item["target_id"]
if "subject" in item and "source" not in item:
item["source"] = item["subject"]
if "object" in item and "target" not in item:
item["target"] = item["object"]
if "source" in item and not isinstance(item["source"], str):
src = item["source"]
item["source"] = getattr(src, "id", getattr(src, "text", str(src)))
@@ -347,6 +354,21 @@ class GraphBuilder:
elif not isinstance(sources, list):
sources = [sources]
# Count input relationships for warning if all are dropped
input_relationships_count = 0
if isinstance(source_dict, dict):
rels = source_dict.get("relationships", [])
if isinstance(rels, list):
input_relationships_count += len(rels)
elif rels is not None:
input_relationships_count += 1
if explicit_relationships:
for rel_item in explicit_relationships:
if isinstance(rel_item, list):
input_relationships_count += len(rel_item)
else:
input_relationships_count += 1
# Track graph building
build_start_time = time.time()
@@ -468,11 +490,12 @@ class GraphBuilder:
pipeline_id=pipeline_id,
)
# Check if relationships are already in dictionary format
sample_rel = relationships_list[0] if relationships_list else None
is_dict_format = isinstance(sample_rel, dict) and (
"source" in sample_rel and "target" in sample_rel
) and not hasattr(sample_rel, "__dict__") # Ensure it's not a class instance
("source" in sample_rel and "target" in sample_rel)
or ("source_id" in sample_rel and "target_id" in sample_rel)
or ("subject" in sample_rel and "object" in sample_rel)
) and not hasattr(sample_rel, "__dict__")
if is_dict_format:
# Fast path: directly append dictionaries after normalizing source/target
@@ -481,8 +504,15 @@ class GraphBuilder:
batch = relationships_list[i:i + batch_size]
for item in batch:
if isinstance(item, dict):
# Normalize source/target if they are objects
rel_dict = item.copy()
if "source_id" in rel_dict and "source" not in rel_dict:
rel_dict["source"] = rel_dict["source_id"]
if "target_id" in rel_dict and "target" not in rel_dict:
rel_dict["target"] = rel_dict["target_id"]
if "subject" in rel_dict and "source" not in rel_dict:
rel_dict["source"] = rel_dict["subject"]
if "object" in rel_dict and "target" not in rel_dict:
rel_dict["target"] = rel_dict["object"]
if "source" in rel_dict and not isinstance(rel_dict["source"], str):
src = rel_dict["source"]
rel_dict["source"] = getattr(src, "id", getattr(src, "text", str(src)))
@@ -571,6 +601,14 @@ class GraphBuilder:
f"Entity resolution complete: {len(all_entities)} -> {len(resolved_entities)} unique entities"
)
if input_relationships_count > 0 and len(all_relationships) == 0:
warning_msg = (
f"All relationships were dropped during graph building: "
f"{input_relationships_count} input relationships, 0 in final graph"
)
self.logger.warning(warning_msg)
print(f"Warning: {warning_msg}")
# Build graph structure
print("Building graph structure...")
structure_start = time.time()
@@ -682,6 +720,14 @@ class GraphBuilder:
)
raise
def build_single_source(
self,
kg_data: Dict[str, Any],
pipeline_id: Optional[str] = None,
**options,
) -> Dict[str, Any]:
return self.build(kg_data, pipeline_id=pipeline_id, **options)
def add_temporal_edge(
self,
graph,
+8 -1
View File
@@ -109,12 +109,14 @@ Convenience Functions:
- create_associative_class: Associative class creation wrapper
- get_ontology_method: Get ontology method by name
- list_available_methods: List registered methods
- ingest_ontology: Ingest ontology from file or directory
Example Usage:
>>> from semantica.ontology import generate_ontology, infer_classes, OntologyGenerator
>>> from semantica.ontology import generate_ontology, infer_classes, OntologyGenerator, ingest_ontology
>>> # Using convenience functions
>>> ontology = generate_ontology({"entities": [...], "relationships": [...]}, method="default")
>>> classes = infer_classes(entities, method="default")
>>> data = ingest_ontology("ontology.ttl")
>>> # Using classes directly
>>> from semantica.ontology import OntologyGenerator, ClassInferrer, PropertyGenerator
>>> generator = OntologyGenerator(base_uri="https://example.org/ontology/")
@@ -155,6 +157,8 @@ from .registry import MethodRegistry, method_registry
from .requirements_spec import RequirementsSpec, RequirementsSpecManager
from .reuse_manager import ReuseDecision, ReuseManager
from .version_manager import OntologyVersion, VersionManager
from semantica.ingest import OntologyData, OntologyIngestor
from .methods import ingest_ontology
__all__ = [
# Main generators
@@ -200,4 +204,7 @@ __all__ = [
# Configuration
"OntologyConfig",
"ontology_config",
"ingest_ontology",
"OntologyData",
"OntologyIngestor",
]
+26 -3
View File
@@ -111,15 +111,19 @@ Main Functions:
- create_associative_class: Associative class creation wrapper
- get_ontology_method: Get ontology method by name
- list_available_methods: List registered methods
- ingest_ontology: Ingest ontology from file or directory (via semantica.ingest)
Example Usage:
>>> from semantica.ontology.methods import generate_ontology, infer_classes
>>> from semantica.ontology.methods import generate_ontology, infer_classes, ingest_ontology
>>> ontology = generate_ontology({"entities": [...], "relationships": [...]}, method="default")
>>> classes = infer_classes(entities, method="default")
>>> data = ingest_ontology("ontology.ttl")
"""
from typing import Any, Callable, Dict, List, Optional
from typing import Any, Callable, Dict, List, Optional, Union
from pathlib import Path
from semantica.ingest import ingest_ontology as _ingest_ontology, OntologyData
from .registry import method_registry
@@ -172,4 +176,23 @@ def list_available_methods(task: Optional[str] = None) -> Dict[str, List[str]]:
return method_registry.list_all(task)
pass
def ingest_ontology(
source: Union[str, Path, List[Union[str, Path]]],
method: str = "file",
**kwargs
) -> Union[OntologyData, List[OntologyData]]:
"""
Ingest ontology from source.
This is a convenience wrapper around semantica.ingest.ingest_ontology.
Args:
source: Ontology file path, directory path, or list of paths
method: Ingestion method (default: "file")
**kwargs: Additional options
Returns:
OntologyData or List[OntologyData]
"""
return _ingest_ontology(source, method=method, **kwargs)
+55
View File
@@ -42,6 +42,18 @@ classes = inferrer.infer_classes(entities, build_hierarchy=True)
properties = prop_gen.infer_properties(entities, relationships, classes)
```
### Ingesting Ontologies
```python
from semantica.ingest import OntologyIngestor
# Create ingestor
ingestor = OntologyIngestor()
# Ingest ontology
ontology_data = ingestor.ingest_ontology("ontology.ttl")
```
## Ontology Generation
### Basic Ontology Generation
@@ -115,6 +127,49 @@ ontology = engine.from_data(
)
```
## Ontology Ingestion
### Basic Ingestion
Ingest existing ontologies from files (Turtle, RDF/XML, JSON-LD, etc.) into `OntologyData` objects.
```python
from semantica.ontology import ingest_ontology
# Ingest a single file
ontology_data = ingest_ontology("path/to/ontology.ttl")
print(f"Source: {ontology_data.source_path}")
print(f"Format: {ontology_data.format}")
print(f"Data keys: {ontology_data.data.keys()}")
```
### Ingesting Directories
Ingest all ontology files in a directory recursively.
```python
from semantica.ontology import ingest_ontology
# Ingest a directory
ontologies = ingest_ontology("path/to/ontologies_dir/")
for ont in ontologies:
print(f"Ingested: {ont.source_path} ({ont.format})")
```
### Unified Ingestion Interface
You can also use the unified `semantica.ingest` interface.
```python
from semantica.ingest import ingest
# Ingest as "ontology" source type
result = ingest("path/to/ontology.ttl", source_type="ontology")
ontology_data = result["ontology"]
```
## Class Inference
### Basic Class Inference
+87
View File
@@ -363,3 +363,90 @@ class ReuseManager:
def list_known_ontologies(self) -> List[str]:
"""List known ontology URIs."""
return list(self.known_ontologies.keys())
def merge_ontology_data(
self, target: Dict[str, Any], source: Dict[str, Any], **options
) -> Dict[str, Any]:
"""
Merge source ontology data into target ontology.
Merges classes, properties, and metadata from source to target.
Handles deduplication based on URI and name.
Args:
target: Target ontology dictionary (modified in-place)
source: Source ontology dictionary
**options: Merge options:
- overwrite: Whether to overwrite existing elements (default: False)
- merge_metadata: Whether to merge metadata (default: True)
Returns:
Merged target ontology
"""
tracking_id = self.progress_tracker.start_tracking(
module="ontology",
submodule="ReuseManager",
message=f"Merging ontology {source.get('name', 'unknown')} into {target.get('name', 'unknown')}",
)
try:
overwrite = options.get("overwrite", False)
# Helper to merge lists of dicts (classes/properties)
def merge_lists(target_list, source_list, key_field="uri"):
existing_keys = {item.get(key_field): i for i, item in enumerate(target_list) if item.get(key_field)}
for item in source_list:
key = item.get(key_field)
if not key:
# Fallback to name if URI missing
key = item.get("name")
if key in existing_keys:
if overwrite:
target_list[existing_keys[key]] = item
else:
target_list.append(item)
if key:
existing_keys[key] = len(target_list) - 1
# Merge Classes
if "classes" in source:
if "classes" not in target:
target["classes"] = []
merge_lists(target["classes"], source["classes"])
# Merge Properties
if "properties" in source:
if "properties" not in target:
target["properties"] = []
merge_lists(target["properties"], source["properties"])
# Merge Metadata
if options.get("merge_metadata", True) and "metadata" in source:
if "metadata" not in target:
target["metadata"] = {}
# Update with source metadata, preserving target's specific fields if needed
# Here we just update
target["metadata"].update(source["metadata"])
# Merge Imports
if "imports" in source:
if "imports" not in target:
target["imports"] = []
for imp in source["imports"]:
if imp not in target["imports"]:
target["imports"].append(imp)
self.progress_tracker.stop_tracking(
tracking_id,
status="completed",
message=f"Merged ontology data successfully",
)
return target
except Exception as e:
self.progress_tracker.stop_tracking(
tracking_id, status="failed", message=str(e)
)
raise
+6 -2
View File
@@ -32,12 +32,14 @@ License: MIT
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Callable, Dict, List, Optional, Union
from typing import Any, Callable, Dict, List, Optional, Union, TYPE_CHECKING
from ..utils.exceptions import ProcessingError, ValidationError
from ..utils.logging import get_logger
from ..utils.progress_tracker import get_progress_tracker
from .pipeline_validator import PipelineValidator
if TYPE_CHECKING:
from .pipeline_validator import PipelineValidator
class StepStatus(Enum):
@@ -104,6 +106,8 @@ class PipelineBuilder:
if not self.progress_tracker.enabled:
self.progress_tracker.enabled = True
from .pipeline_validator import PipelineValidator
self.validator = PipelineValidator(**self.config)
self.steps: List[PipelineStep] = []
self.step_registry: Dict[str, Callable] = {}
+184
View File
@@ -0,0 +1,184 @@
"""
Result Caching Module
This module provides caching mechanisms for extraction results to avoid redundant
computations and API calls. It implements an LRU (Least Recently Used) cache
with Time-To-Live (TTL) support.
Key Features:
- LRU Caching: Evicts least recently used items when cache is full
- TTL Support: Expires items after a configurable duration
- Namespaced Caching: Separate caches for entities, relations, and triplets
- Hash-based Keys: Uses stable hashing for text and parameters
Classes:
- ExtractionCache: Main cache manager
- CacheItem: Container for cached data with metadata
Author: Semantica Contributors
License: MIT
"""
import time
import hashlib
import json
from collections import OrderedDict
from typing import Any, Dict, Optional, Union, List
from threading import Lock
from ..utils.logging import get_logger
class CacheItem:
"""Container for cached data."""
def __init__(self, value: Any, ttl: Optional[int] = None):
self.value = value
self.timestamp = time.time()
self.ttl = ttl
def is_expired(self) -> bool:
"""Check if item has expired."""
if self.ttl is None:
return False
return time.time() - self.timestamp > self.ttl
class ExtractionCache:
"""
LRU Cache for extraction results.
Thread-safe implementation.
"""
def __init__(self, max_size: int = 1000, ttl: int = 3600):
"""
Initialize the cache.
Args:
max_size: Maximum number of items to store per namespace
ttl: Time to live in seconds (default 1 hour)
"""
self.max_size = max_size
self.ttl = ttl
self._caches: Dict[str, OrderedDict] = {
"entities": OrderedDict(),
"relations": OrderedDict(),
"triplets": OrderedDict()
}
self._locks: Dict[str, Lock] = {
"entities": Lock(),
"relations": Lock(),
"triplets": Lock()
}
self.logger = get_logger("extraction_cache")
self.enabled = True
def _generate_key(self, text: str, **params) -> str:
"""
Generate a stable cache key based on text and parameters.
Note: Sensitive parameters like 'api_key' are excluded from the cache key
to prevent security risks and ensure cache sharing where appropriate.
"""
# Filter out sensitive keys
sensitive_keys = {'api_key', 'token', 'password', 'secret', 'auth', 'authorization'}
filtered_params = {k: v for k, v in params.items() if k.lower() not in sensitive_keys}
# Create a stable string representation of params
# Sort keys to ensure consistent ordering
param_str = json.dumps(filtered_params, sort_keys=True, default=str)
# Combine text and params
content = f"{text}|{param_str}"
# Return hash (SHA-256 for better security than MD5)
return hashlib.sha256(content.encode('utf-8')).hexdigest()
def get(self, namespace: str, text: str, **params) -> Optional[Any]:
"""
Retrieve item from cache.
Args:
namespace: Cache namespace ("entities", "relations", "triplets")
text: Input text used for extraction
**params: Extraction parameters used
Returns:
Cached result or None if not found/expired
"""
if not self.enabled:
return None
if namespace not in self._caches:
return None
key = self._generate_key(text, **params)
with self._locks[namespace]:
cache = self._caches[namespace]
if key in cache:
item = cache[key]
# Check expiration
if item.is_expired():
del cache[key]
return None
# Move to end (mark as recently used)
cache.move_to_end(key)
return item.value
return None
def set(self, namespace: str, text: str, value: Any, **params) -> None:
"""
Add item to cache.
Args:
namespace: Cache namespace
text: Input text
value: Result to cache
**params: Extraction parameters
"""
if not self.enabled:
return
if namespace not in self._caches:
self.logger.warning(f"Unknown cache namespace: {namespace}")
return
key = self._generate_key(text, **params)
item = CacheItem(value, self.ttl)
with self._locks[namespace]:
cache = self._caches[namespace]
# If key exists, update and move to end
if key in cache:
cache.move_to_end(key)
cache[key] = item
# Evict if full
if len(cache) > self.max_size:
cache.popitem(last=False) # Remove first (least recently used)
def clear(self, namespace: Optional[str] = None):
"""Clear cache(s)."""
if namespace:
if namespace in self._caches:
with self._locks[namespace]:
self._caches[namespace].clear()
else:
for ns in self._caches:
with self._locks[ns]:
self._caches[ns].clear()
def get_stats(self) -> Dict[str, Dict[str, int]]:
"""Get cache statistics."""
stats = {}
for ns, cache in self._caches.items():
stats[ns] = {
"size": len(cache),
"max_size": self.max_size
}
return stats
# Global cache instance
extraction_cache = ExtractionCache()
+76 -1
View File
@@ -40,8 +40,9 @@ License: MIT
"""
import os
import multiprocessing
from pathlib import Path
from typing import Dict, Optional
from typing import Dict, Optional, Any
from ..utils.logging import get_logger
@@ -53,9 +54,23 @@ class Config:
"""Initialize configuration manager."""
self.logger = get_logger("config")
self._configs: Dict[str, Dict] = {}
# Default optimization settings
self._configs["optimization"] = {
"enable_cache": True,
"cache_size": 1000,
"max_workers": 8,
"enable_batching": True,
"batch_size": 10,
"max_tokens_per_batch": 2000
}
self._load_config_file(config_file)
self._load_env_vars()
def get_optimization_config(self) -> Dict:
"""Get optimization configuration."""
return self._configs.get("optimization", {})
def _load_config_file(self, config_file: Optional[str]):
"""Load configuration from file."""
if config_file and Path(config_file).exists():
@@ -114,6 +129,66 @@ class Config:
return self._configs[provider].get("api_key")
return os.getenv(f"{provider.upper()}_API_KEY")
def get(self, key: str, default: Any = None) -> Any:
"""
Get configuration value by key.
Searches in top-level configs and optimization settings.
"""
# 1. Check top-level keys
if key in self._configs:
return self._configs[key]
# 2. Check optimization settings (common keys)
if "optimization" in self._configs and key in self._configs["optimization"]:
return self._configs["optimization"][key]
# 3. Handle specific mapping for optimization keys
# Map cache_enabled -> enable_cache if needed
if key == "cache_enabled":
return self._configs.get("optimization", {}).get("enable_cache", default)
return default
# Global config instance
config = Config()
def resolve_max_workers(
explicit: Optional[int] = None,
local_config: Optional[Dict[str, Any]] = None,
methods: Optional[Any] = None,
) -> int:
def to_int(val: Any, default: int) -> int:
try:
return int(val)
except Exception:
return default
if isinstance(methods, str):
normalized_methods = [methods]
elif isinstance(methods, (list, tuple, set)):
normalized_methods = [m for m in methods if isinstance(m, str)]
else:
normalized_methods = []
if explicit is not None:
value = to_int(explicit, 1)
elif local_config and "max_workers" in local_config:
value = to_int(local_config.get("max_workers", 1), 1)
else:
value = to_int(config.get("max_workers", 5), 5)
if "ml" in normalized_methods and explicit is None and not (local_config and "max_workers" in local_config):
value = 1
if value < 1:
value = 1
cpu_count = multiprocessing.cpu_count() or 1
if value > cpu_count:
value = cpu_count
if value > 32:
value = 32
return value
@@ -240,6 +240,12 @@ class CoreferenceResolver:
self.progress_tracker.stop_tracking(
tracking_id, status="failed", message=str(e)
)
verbose_mode = options.get("verbose", False) or self.config.get("verbose", False)
if verbose_mode:
import sys
print(f" [CoreferenceResolver] ERROR: Resolution failed: {e}", flush=True, file=sys.stderr)
import traceback
traceback.print_exc(file=sys.stderr)
raise
def resolve(
+118 -73
View File
@@ -85,68 +85,59 @@ class Event:
class EventDetector:
"""Event detection and extraction handler."""
def __init__(
self,
event_types: Optional[List[str]] = None,
extract_participants: bool = True,
extract_location: bool = True,
extract_time: bool = True,
method: Union[str, List[str]] = None,
config=None,
**kwargs
):
def __init__(self, method: str = "llm", **config):
"""
Initialize event detector.
Args:
event_types: Specific event types to detect (e.g., ["launch", "acquisition"])
extract_participants: Whether to extract event participants
extract_location: Whether to extract event locations
extract_time: Whether to extract temporal information
method: Extraction method(s) for underlying NER/relation extractors.
Can be passed to ner_method and relation_method in config.
config: Legacy config dict (deprecated, use kwargs)
**kwargs: Configuration options:
- ner_method: Method for NER extraction (if entities need to be extracted)
- relation_method: Method for relation extraction (if relations need to be extracted)
- Other options passed to sub-components
method: Extraction method ("llm", "pattern")
**config: Configuration options
"""
self.logger = get_logger("event_detector")
self.config = config or {}
self.config.update(kwargs)
self.config = config
self.method = method
self.progress_tracker = get_progress_tracker()
# Ensure progress tracker is enabled
if not self.progress_tracker.enabled:
self.progress_tracker.enabled = True
# Store parameters
self.event_types_filter = event_types
self.extract_participants = extract_participants
self.extract_location = extract_location
self.extract_time = extract_time
# Initialize components
self.event_classifier = EventClassifier(**config)
self.temporal_processor = TemporalEventProcessor(**config)
# Configure extraction options
self.extract_participants = config.get("extract_participants", True)
self.extract_location = config.get("extract_location", True)
self.extract_time = config.get("extract_time", True)
self.event_types_filter = config.get("event_types", [])
# Define event patterns
self.event_patterns = {
"acquisition": r"\b(acquired|acquisition|buying|bought|merger|merged)\b",
"partnership": r"\b(partnered|partnership|collaborate|collaboration)\b",
"launch": r"\b(launch|launched|releasing|released|unveil|unveiled)\b",
"investment": r"\b(invest|invested|investment|funding|raised)\b",
"legal": r"\b(sue|sued|lawsuit|litigation|legal action)\b",
}
# Pre-compile location patterns
self.location_patterns = [
re.compile(r"in\s+([A-Z][a-z]+(?:\s+[A-Z][a-z]+)*)"),
re.compile(r"at\s+([A-Z][a-z]+(?:\s+[A-Z][a-z]+)*)"),
]
# Pre-compile time patterns
self.time_patterns = [
re.compile(r"on\s+([A-Z][a-z]+\s+\d{1,2},?\s+\d{4})"),
re.compile(r"in\s+(\d{4})"),
re.compile(r"(\d{1,2}[/-]\d{1,2}[/-]\d{2,4})"),
]
# Store method for passing to extractors if needed
if method is not None:
self.config["ner_method"] = method
self.config["relation_method"] = method
self.event_classifier = EventClassifier(**self.config.get("classifier", {}))
self.temporal_processor = TemporalEventProcessor(
**self.config.get("temporal", {})
)
self.relationship_extractor = EventRelationshipExtractor(
**self.config.get("relationship", {})
)
# Event patterns
self.event_patterns = {
"founded": r"founded|created|established",
"acquired": r"acquired|bought|purchased",
"launched": r"launched|released|introduced",
"announced": r"announced|declared|stated",
"meeting": r"met|meeting|conference|summit",
}
def extract(
self,
text: Union[str, List[str], List[Dict[str, Any]]],
@@ -175,9 +166,10 @@ class EventDetector:
)
try:
results = []
results = [None] * len(text) # Pre-allocate to maintain order
total_items = len(text)
total_events_count = 0
processed_count = 0
# Determine update interval
if total_items <= 10:
@@ -193,33 +185,77 @@ class EventDetector:
message=f"Starting batch detection... 0/{total_items} (remaining: {total_items})"
)
for idx, item in enumerate(text):
# Prepare arguments for single item
doc_text = item["content"] if isinstance(item, dict) and "content" in item else str(item)
from .config import resolve_max_workers
max_workers = resolve_max_workers(
explicit=kwargs.get("max_workers"),
local_config=self.config,
methods=[self.config.get("ner_method"), self.config.get("relation_method"), self.config.get("method")],
)
def process_item(idx, item):
try:
# Prepare arguments for single item
doc_text = item["content"] if isinstance(item, dict) and "content" in item else str(item)
# Detect
events = self.detect_events(doc_text, **kwargs)
# Add provenance metadata
for event in events:
if event.metadata is None:
event.metadata = {}
event.metadata["batch_index"] = idx
if isinstance(item, dict) and "id" in item:
event.metadata["document_id"] = item["id"]
return idx, events
except Exception as e:
self.logger.error(f"Error processing item {idx}: {e}")
# Return empty list on failure to continue processing
return idx, []
if max_workers > 1:
import concurrent.futures
# Detect
events = self.detect_events(doc_text, **kwargs)
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
# Submit tasks
future_to_idx = {}
for idx, item in enumerate(text):
future = executor.submit(process_item, idx, item)
future_to_idx[future] = idx
for future in concurrent.futures.as_completed(future_to_idx):
idx, events = future.result()
results[idx] = events
total_events_count += len(events)
processed_count += 1
# Update progress
if processed_count % update_interval == 0 or processed_count == total_items:
remaining = total_items - processed_count
self.progress_tracker.update_progress(
tracking_id,
processed=processed_count,
total=total_items,
message=f"Processing... {processed_count}/{total_items} (remaining: {remaining}) - Detected {total_events_count} events"
)
else:
# Sequential processing
for idx, item in enumerate(text):
_, events = process_item(idx, item)
results[idx] = events
total_events_count += len(events)
processed_count += 1
# Add provenance metadata
for event in events:
if event.metadata is None:
event.metadata = {}
event.metadata["batch_index"] = idx
if isinstance(item, dict) and "id" in item:
event.metadata["document_id"] = item["id"]
results.append(events)
total_events_count += len(events)
# Update progress
if (idx + 1) % update_interval == 0 or (idx + 1) == total_items:
remaining = total_items - (idx + 1)
self.progress_tracker.update_progress(
tracking_id,
processed=idx + 1,
total=total_items,
message=f"Processing... {idx + 1}/{total_items} (remaining: {remaining}) - Detected {total_events_count} events"
)
# Update progress
if processed_count % update_interval == 0 or processed_count == total_items:
remaining = total_items - processed_count
self.progress_tracker.update_progress(
tracking_id,
processed=processed_count,
total=total_items,
message=f"Processing... {processed_count}/{total_items} (remaining: {remaining}) - Detected {total_events_count} events"
)
self.progress_tracker.stop_tracking(
tracking_id,
@@ -238,17 +274,26 @@ class EventDetector:
# Single item
return self.detect_events(text, **kwargs)
def detect_events(self, text: str, **options) -> List[Event]:
def detect_events(
self,
text: Union[str, List[str], List[Dict[str, Any]]],
pipeline_id: Optional[str] = None,
**options,
) -> Union[List[Event], List[List[Event]]]:
"""
Detect events in text content.
Args:
text: Input text
pipeline_id: Optional pipeline ID for progress tracking (batch mode)
**options: Detection options
Returns:
list: List of detected events
"""
if isinstance(text, list):
return self.extract(text, pipeline_id=pipeline_id, **options)
tracking_id = self.progress_tracker.start_tracking(
module="semantic_extract",
submodule="EventDetector",
+6 -1
View File
@@ -108,7 +108,12 @@ class LLMExtraction:
# Initialize provider using new system
try:
self.provider = create_provider(provider, **config)
# Sanitize config: remove api_key if it's None/empty to allow fallback
provider_config = config.copy()
if "api_key" in provider_config and not provider_config["api_key"]:
del provider_config["api_key"]
self.provider = create_provider(provider, **provider_config)
except Exception as e:
self.logger.warning(f"Failed to initialize {provider} provider: {e}")
self.provider = None
File diff suppressed because it is too large Load Diff
+87 -31
View File
@@ -178,9 +178,11 @@ class NERExtractor:
)
try:
results = []
results = [None] * len(text)
total_items = len(text)
total_entities_count = 0
processed_count = 0
# Update more frequently: every 1% or at least every 10 items, but always update for small datasets
if total_items <= 10:
update_interval = 1 # Update every item for small datasets
@@ -188,15 +190,22 @@ class NERExtractor:
update_interval = max(1, min(10, total_items // 100))
# Initial progress update - ALWAYS show this
remaining = total_items
self.progress_tracker.update_progress(
tracking_id,
processed=0,
total=total_items,
message=f"Starting batch extraction... 0/{total_items} (remaining: {remaining})"
message=f"Starting batch extraction... 0/{total_items}"
)
for idx, item in enumerate(text, 1):
from .config import resolve_max_workers
max_workers = resolve_max_workers(
explicit=kwargs.get("max_workers"),
local_config=self.config,
methods=self.method,
)
# Helper function for single item processing
def process_item(idx, item):
try:
current_entities = []
if isinstance(item, dict) and "content" in item:
@@ -214,30 +223,69 @@ class NERExtractor:
for ent in current_entities:
if ent.metadata is None:
ent.metadata = {}
ent.metadata["batch_index"] = idx - 1
ent.metadata["batch_index"] = idx
if isinstance(item, dict) and "id" in item:
ent.metadata["document_id"] = item["id"]
results.append(current_entities)
total_entities_count += len(current_entities)
except Exception:
results.append([])
return idx, current_entities
except Exception as e:
self.logger.warning(f"Failed to process item {idx}: {e}")
return idx, []
if max_workers > 1:
import concurrent.futures
remaining = total_items - idx
# Update progress: always update for small datasets, or at intervals for large ones
should_update = (
idx % update_interval == 0 or
idx == total_items or
idx == 1 or
total_items <= 10 # Always update for small datasets
)
if should_update:
self.progress_tracker.update_progress(
tracking_id,
processed=idx,
total=total_items,
message=f"Processing documents... {idx}/{total_items} (remaining: {remaining}) - Extracted {total_entities_count} entities so far"
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
# Submit all tasks
future_to_idx = {
executor.submit(process_item, idx, item): idx
for idx, item in enumerate(text)
}
for future in concurrent.futures.as_completed(future_to_idx):
idx, entities = future.result()
results[idx] = entities
total_entities_count += len(entities)
processed_count += 1
# Update progress
should_update = (
processed_count % update_interval == 0 or
processed_count == total_items or
processed_count == 1 or
total_items <= 10
)
if should_update:
remaining = total_items - processed_count
self.progress_tracker.update_progress(
tracking_id,
processed=processed_count,
total=total_items,
message=f"Processing documents... {processed_count}/{total_items} (remaining: {remaining}) - Extracted {total_entities_count} entities so far"
)
else:
# Sequential processing
for idx, item in enumerate(text):
_, entities = process_item(idx, item)
results[idx] = entities
total_entities_count += len(entities)
processed_count += 1
# Update progress
should_update = (
processed_count % update_interval == 0 or
processed_count == total_items or
processed_count == 1 or
total_items <= 10
)
if should_update:
remaining = total_items - processed_count
self.progress_tracker.update_progress(
tracking_id,
processed=processed_count,
total=total_items,
message=f"Processing documents... {processed_count}/{total_items} (remaining: {remaining}) - Extracted {total_entities_count} entities so far"
)
self.progress_tracker.stop_tracking(
tracking_id,
@@ -253,12 +301,18 @@ class NERExtractor:
else:
return self.extract_entities(text, **kwargs)
def extract_entities(self, text: str, **options) -> List[Entity]:
def extract_entities(
self,
text: Union[str, List[Dict[str, Any]], List[str]],
pipeline_id: Optional[str] = None,
**options,
) -> Union[List[Entity], List[List[Entity]]]:
"""
Extract named entities from text.
Args:
text: Input text
pipeline_id: Optional pipeline ID for progress tracking (batch mode)
**options: Extraction options:
- entity_types: Filter by entity types (list)
- min_confidence: Minimum confidence threshold
@@ -267,6 +321,9 @@ class NERExtractor:
Returns:
list: List of extracted entities
"""
if isinstance(text, list):
return self.extract(text, pipeline_id=pipeline_id, **options)
tracking_id = self.progress_tracker.start_tracking(
module="semantic_extract",
submodule="NERExtractor",
@@ -318,14 +375,13 @@ class NERExtractor:
method_options["model"] = all_options.get(
"llm_model", all_options.get("model")
)
# Pass api_key if provided (needed for all providers)
if "api_key" in all_options:
method_options["api_key"] = all_options["api_key"]
elif "api_key" not in method_options:
# Try to get from environment as fallback
# Ensure api_key is populated: check explicitly provided or fallback to env
current_key = method_options.get("api_key")
if not current_key:
# Not found or empty/None, try environment
import os
provider = method_options.get("provider", "openai")
env_key = f"{provider.upper()}_API_KEY"
provider_name = method_options.get("provider", "openai")
env_key = f"{provider_name.upper()}_API_KEY"
api_key = os.getenv(env_key)
if api_key:
method_options["api_key"] = api_key
+482 -143
View File
@@ -243,68 +243,136 @@ class BaseProvider:
mode = instructor.Mode.TOOLS # Default mode
if provider_name == "OpenAIProvider" and self.client:
client = instructor.from_openai(self.client)
elif provider_name == "AnthropicProvider" and self.client:
client = instructor.from_anthropic(self.client)
elif provider_name == "GeminiProvider" and self.client:
client = instructor.from_gemini(
self.client,
mode=instructor.Mode.GEMINI_JSON
)
elif provider_name == "GroqProvider" and self.client:
# Try using from_groq if available (newer instructor versions)
if hasattr(instructor, "from_groq"):
client = instructor.from_groq(self.client, mode=instructor.Mode.JSON)
if hasattr(instructor, "from_provider"):
try:
client = instructor.from_provider(
provider=f"openai/{kwargs.get('model', self.model)}",
api_key=self.api_key
)
except Exception:
client = instructor.from_openai(self.client)
else:
# Fallback: Create OpenAI client pointing to Groq
# This avoids the "Client should be an instance of openai.OpenAI" warning
client = instructor.from_openai(self.client)
elif provider_name == "AnthropicProvider" and self.client:
if hasattr(instructor, "from_provider"):
try:
client = instructor.from_provider(
provider=f"anthropic/{kwargs.get('model', self.model)}",
api_key=self.api_key
)
except Exception:
client = instructor.from_anthropic(self.client)
else:
client = instructor.from_anthropic(self.client)
elif provider_name == "GeminiProvider" and self.client:
if hasattr(instructor, "from_provider"):
try:
client = instructor.from_provider(
provider=f"gemini/{kwargs.get('model', self.model)}",
api_key=self.api_key
)
except Exception:
client = instructor.from_gemini(
self.client,
mode=instructor.Mode.GEMINI_JSON
)
else:
client = instructor.from_gemini(
self.client,
mode=instructor.Mode.GEMINI_JSON
)
elif provider_name == "GroqProvider" and self.client:
# Try using from_provider which is recommended for Groq in latest instructor
if hasattr(instructor, "from_provider"):
try:
client = instructor.from_provider(
provider=f"groq/{kwargs.get('model', self.model)}",
api_key=self.api_key
)
except Exception:
client = None
if not client:
# Try using from_groq if available (newer instructor versions)
if hasattr(instructor, "from_groq"):
client = instructor.from_groq(self.client, mode=instructor.Mode.JSON)
else:
# Fallback: Create OpenAI client pointing to Groq
# This avoids the "Client should be an instance of openai.OpenAI" warning
try:
from openai import OpenAI
# Fix: Use self.api_key instead of self.client.api_key
groq_client = OpenAI(
base_url="https://api.groq.com/openai/v1",
api_key=self.api_key,
)
client = instructor.from_openai(groq_client, mode=instructor.Mode.JSON)
except Exception:
# Last resort: try passing the groq client directly
client = instructor.from_openai(self.client, mode=instructor.Mode.JSON)
elif provider_name == "OllamaProvider":
# Try from_provider for Ollama if available
if hasattr(instructor, "from_provider"):
try:
client = instructor.from_provider(
provider=f"ollama/{kwargs.get('model', self.model)}",
)
except Exception:
client = None
if not client:
# Create OpenAI-compatible client for Ollama
try:
from openai import OpenAI
groq_client = OpenAI(
base_url="https://api.groq.com/openai/v1",
api_key=self.client.api_key,
# Ollama typically runs on localhost:11434/v1
base_url = getattr(self, "base_url", "http://localhost:11434")
if not base_url.endswith("/v1"):
base_url = f"{base_url.rstrip('/')}/v1"
ollama_client = OpenAI(
base_url=base_url,
api_key="ollama", # required but unused
)
client = instructor.from_openai(groq_client, mode=instructor.Mode.JSON)
except Exception:
# Last resort: try passing the groq client directly
client = instructor.from_openai(self.client, mode=instructor.Mode.JSON)
elif provider_name == "OllamaProvider":
# Create OpenAI-compatible client for Ollama
try:
from openai import OpenAI
# Ollama typically runs on localhost:11434/v1
base_url = getattr(self, "base_url", "http://localhost:11434")
if not base_url.endswith("/v1"):
base_url = f"{base_url.rstrip('/')}/v1"
ollama_client = OpenAI(
base_url=base_url,
api_key="ollama", # required but unused
)
client = instructor.from_openai(ollama_client, mode=instructor.Mode.JSON)
except ImportError:
pass
client = instructor.from_openai(ollama_client, mode=instructor.Mode.JSON)
except ImportError:
pass
elif provider_name == "DeepSeekProvider" and self.client:
# DeepSeek is OpenAI compatible
# We need to wrap the underlying client if it exposes the OpenAI interface
# or create a new OpenAI client if self.client is a deepseek.Client (which might be just a wrapper)
# Assuming deepseek.Client is compatible or we can use OpenAI client
try:
# DeepSeek usually works with standard OpenAI client
# If self.client is deepseek.Client, check if we can wrap it
# Otherwise create a new OpenAI client
from openai import OpenAI
if isinstance(self.client, OpenAI):
client = instructor.from_openai(self.client, mode=instructor.Mode.JSON)
else:
# Try creating fresh client
ds_client = OpenAI(
api_key=self.api_key,
base_url="https://api.deepseek.com"
)
client = instructor.from_openai(ds_client, mode=instructor.Mode.JSON)
except Exception:
pass
# Try from_provider for DeepSeek
if hasattr(instructor, "from_provider"):
try:
client = instructor.from_provider(
provider=f"deepseek/{kwargs.get('model', self.model)}",
api_key=self.api_key
)
except Exception:
client = None
if not client:
# DeepSeek is OpenAI compatible
try:
from openai import OpenAI
if isinstance(self.client, OpenAI):
client = instructor.from_openai(self.client, mode=instructor.Mode.JSON)
else:
# Try creating fresh client
ds_client = OpenAI(
api_key=self.api_key,
base_url="https://api.deepseek.com"
)
client = instructor.from_openai(ds_client, mode=instructor.Mode.JSON)
except Exception:
pass
# Global LiteLLM support - if litellm is passed in kwargs or config
if not client and (kwargs.get("litellm") or self.config.get("litellm")):
if hasattr(instructor, "from_provider"):
try:
# Format for litellm in instructor is litellm/model_name
provider_model = kwargs.get("model", self.model)
litellm_provider = f"litellm/{provider_model}"
client = instructor.from_provider(litellm_provider, api_key=self.api_key)
except Exception:
pass
if client:
# Map generate arguments to client arguments
@@ -318,11 +386,24 @@ class BaseProvider:
"temperature": kwargs.get("temperature", 0.1), # Low temp for structured
}
verbose_mode = kwargs.get("verbose", False)
if verbose_mode:
import sys
print(f" [BaseProvider.generate_typed] Using instructor via {provider_name}. Client: {type(client)}", flush=True, file=sys.stdout)
# Pass through other common parameters
for param in ["max_tokens", "max_completion_tokens", "top_p", "frequency_penalty", "presence_penalty", "seed", "stop", "logit_bias", "user", "top_k"]:
if param in kwargs:
create_kwargs[param] = kwargs[param]
# Add provider-specific params
if provider_name == "GroqProvider":
create_kwargs["response_format"] = {"type": "json_object"}
response = client.chat.completions.create(**create_kwargs)
if verbose_mode:
import sys
print(f" [BaseProvider.generate_typed] Typed response received via instructor ({provider_name}).", flush=True, file=sys.stdout)
return response
except Exception as e:
self.logger.warning(f"Instructor generation failed ({e}), falling back to manual repair loop.")
@@ -464,11 +545,24 @@ class OpenAIProvider(BaseProvider):
"OpenAI client not initialized. Set OPENAI_API_KEY or pass api_key."
)
response = self.client.chat.completions.create(
model=kwargs.get("model", self.model),
messages=[{"role": "user", "content": prompt}],
temperature=kwargs.get("temperature", 0.3),
)
create_kwargs = {
"model": kwargs.get("model", self.model),
"messages": [{"role": "user", "content": prompt}],
"temperature": kwargs.get("temperature", 0.3),
}
# Support max_tokens and max_completion_tokens (for o1 models)
if "max_completion_tokens" in kwargs:
create_kwargs["max_completion_tokens"] = kwargs["max_completion_tokens"]
elif "max_tokens" in kwargs:
create_kwargs["max_tokens"] = kwargs["max_tokens"]
# Pass through other common parameters
for param in ["top_p", "frequency_penalty", "presence_penalty", "seed", "stop", "logit_bias", "user"]:
if param in kwargs:
create_kwargs[param] = kwargs[param]
response = self.client.chat.completions.create(**create_kwargs)
return response.choices[0].message.content
def generate_structured(self, prompt: str, **kwargs) -> dict:
@@ -476,12 +570,25 @@ class OpenAIProvider(BaseProvider):
if not self.client:
raise ProcessingError("OpenAI client not initialized.")
response = self.client.chat.completions.create(
model=kwargs.get("model", self.model),
messages=[{"role": "user", "content": prompt}],
response_format={"type": "json_object"},
temperature=kwargs.get("temperature", 0.3),
)
create_kwargs = {
"model": kwargs.get("model", self.model),
"messages": [{"role": "user", "content": prompt}],
"response_format": {"type": "json_object"},
"temperature": kwargs.get("temperature", 0.3),
}
# Support max_tokens and max_completion_tokens
if "max_completion_tokens" in kwargs:
create_kwargs["max_completion_tokens"] = kwargs["max_completion_tokens"]
elif "max_tokens" in kwargs:
create_kwargs["max_tokens"] = kwargs["max_tokens"]
# Pass through other common parameters
for param in ["top_p", "frequency_penalty", "presence_penalty", "seed", "stop", "logit_bias", "user"]:
if param in kwargs:
create_kwargs[param] = kwargs[param]
response = self.client.chat.completions.create(**create_kwargs)
try:
return self._parse_json(response.choices[0].message.content)
except Exception as e:
@@ -499,26 +606,40 @@ class GeminiProvider(BaseProvider):
self.api_key = api_key or config.get_api_key("gemini")
self.model = model
self.client = None
self._use_new_genai = False
self._init_client()
def _init_client(self):
"""Initialize Gemini client."""
try:
import google.generativeai as genai
from google import genai as new_genai
if self.api_key:
genai.configure(api_key=self.api_key)
self.client = genai.GenerativeModel(self.model)
except (ImportError, OSError):
self.client = new_genai.Client(api_key=self.api_key)
self._use_new_genai = True
return
except Exception:
pass
try:
import google.generativeai as old_genai
if self.api_key:
old_genai.configure(api_key=self.api_key)
self.client = old_genai.GenerativeModel(self.model)
self._use_new_genai = False
except Exception:
self.client = None
self.logger.warning(
"google-generativeai library not installed. Install with: pip install semantica[llm-gemini]"
)
self.logger.warning("Gemini SDK not installed. Install with: pip install semantica[llm-gemini]")
def is_available(self) -> bool:
"""Check if provider is available."""
return self.client is not None
def _resp_text(self, resp: Any) -> str:
if hasattr(resp, "text"):
return getattr(resp, "text")
try:
return resp.candidates[0].content.parts[0].text
except Exception:
return str(resp)
def generate(self, prompt: str, **kwargs) -> str:
"""Generate text from prompt."""
if not self.client:
@@ -526,23 +647,46 @@ class GeminiProvider(BaseProvider):
"Gemini client not initialized. Set GEMINI_API_KEY or pass api_key."
)
response = self.client.generate_content(
prompt, generation_config={"temperature": kwargs.get("temperature", 0.3)}
)
return response.text
if self._use_new_genai:
model = kwargs.get("model", self.model)
temperature = kwargs.get("temperature", 0.3)
create_kwargs = {"model": model, "contents": prompt, "config": {"temperature": temperature}}
if "max_tokens" in kwargs:
create_kwargs["config"]["max_output_tokens"] = kwargs["max_tokens"]
for p in ["top_p", "top_k", "stop_sequences", "candidate_count"]:
if p in kwargs:
create_kwargs["config"][p] = kwargs[p]
resp = self.client.models.generate_content(**create_kwargs)
return self._resp_text(resp)
else:
generation_config = {"temperature": kwargs.get("temperature", 0.3)}
if "max_tokens" in kwargs:
generation_config["max_output_tokens"] = kwargs["max_tokens"]
for param in ["top_p", "top_k", "stop_sequences", "candidate_count"]:
if param in kwargs:
generation_config[param] = kwargs[param]
response = self.client.generate_content(prompt, generation_config=generation_config)
return self._resp_text(response)
def generate_structured(self, prompt: str, **kwargs) -> dict:
"""Generate structured output."""
if not self.client:
raise ProcessingError("Gemini client not initialized.")
# Add JSON format instruction to prompt
json_prompt = f"{prompt}\n\nReturn the response as valid JSON only."
response = self.client.generate_content(json_prompt)
try:
return self._parse_json(response.text)
except Exception as e:
raise ProcessingError(f"Failed to parse JSON from Gemini response: {e}")
if self._use_new_genai:
model = kwargs.get("model", self.model)
resp = self.client.models.generate_content(model=model, contents=json_prompt)
try:
return self._parse_json(self._resp_text(resp))
except Exception as e:
raise ProcessingError(f"Failed to parse JSON from Gemini response: {e}")
else:
response = self.client.generate_content(json_prompt)
try:
return self._parse_json(self._resp_text(response))
except Exception as e:
raise ProcessingError(f"Failed to parse JSON from Gemini response: {e}")
class GroqProvider(BaseProvider):
@@ -621,11 +765,32 @@ class GroqProvider(BaseProvider):
"Groq client not initialized. Set GROQ_API_KEY or pass api_key."
)
response = self.client.chat.completions.create(
model=kwargs.get("model", self.model),
messages=[{"role": "user", "content": prompt}],
temperature=kwargs.get("temperature", 0.3),
)
create_kwargs = {
"model": kwargs.get("model", self.model),
"messages": [{"role": "user", "content": prompt}],
"temperature": kwargs.get("temperature", 0.3),
}
# Support max_tokens and max_completion_tokens
if "max_completion_tokens" in kwargs:
create_kwargs["max_completion_tokens"] = kwargs["max_completion_tokens"]
elif "max_tokens" in kwargs:
create_kwargs["max_tokens"] = kwargs["max_tokens"]
# Pass through other common parameters
for param in ["top_p", "frequency_penalty", "presence_penalty", "seed", "stop", "user"]:
if param in kwargs:
create_kwargs[param] = kwargs[param]
verbose_mode = kwargs.get("verbose", False)
if verbose_mode:
import sys
print(f" [GroqProvider.generate] Sending request to Groq API (model: {create_kwargs['model']})...", flush=True, file=sys.stdout)
response = self.client.chat.completions.create(**create_kwargs)
if verbose_mode:
import sys
print(f" [GroqProvider.generate] Response received from Groq.", flush=True, file=sys.stdout)
return response.choices[0].message.content
def generate_structured(self, prompt: str, **kwargs) -> dict:
@@ -638,12 +803,33 @@ class GroqProvider(BaseProvider):
if "json" not in prompt.lower():
json_prompt = f"{prompt}\n\nReturn the response as valid JSON only."
response = self.client.chat.completions.create(
model=kwargs.get("model", self.model),
messages=[{"role": "user", "content": json_prompt}],
temperature=kwargs.get("temperature", 0.3),
response_format={"type": "json_object"},
)
create_kwargs = {
"model": kwargs.get("model", self.model),
"messages": [{"role": "user", "content": json_prompt}],
"temperature": kwargs.get("temperature", 0.3),
"response_format": {"type": "json_object"},
}
# Support max_tokens and max_completion_tokens
if "max_completion_tokens" in kwargs:
create_kwargs["max_completion_tokens"] = kwargs["max_completion_tokens"]
elif "max_tokens" in kwargs:
create_kwargs["max_tokens"] = kwargs["max_tokens"]
# Pass through other common parameters
for param in ["top_p", "frequency_penalty", "presence_penalty", "seed", "stop", "user"]:
if param in kwargs:
create_kwargs[param] = kwargs[param]
verbose_mode = kwargs.get("verbose", False)
if verbose_mode:
import sys
print(f" [GroqProvider.generate_structured] Sending structured request to Groq API (model: {create_kwargs['model']})...", flush=True, file=sys.stdout)
response = self.client.chat.completions.create(**create_kwargs)
if verbose_mode:
import sys
print(f" [GroqProvider.generate_structured] Structured response received from Groq.", flush=True, file=sys.stdout)
try:
return self._parse_json(response.choices[0].message.content)
except Exception as e:
@@ -690,11 +876,23 @@ class AnthropicProvider(BaseProvider):
"Anthropic client not initialized. Set ANTHROPIC_API_KEY or pass api_key."
)
response = self.client.messages.create(
model=kwargs.get("model", self.model),
max_tokens=kwargs.get("max_tokens", 4096),
messages=[{"role": "user", "content": prompt}],
)
# Anthropic requires max_tokens.
# We rely on kwargs, but fallback to 8192 (safe max for newer models) if not provided.
max_tokens = kwargs.get("max_tokens", 8192)
# Prepare arguments
create_kwargs = {
"model": kwargs.get("model", self.model),
"max_tokens": max_tokens,
"messages": [{"role": "user", "content": prompt}],
}
# Pass through other common parameters
for param in ["temperature", "top_p", "top_k", "stop_sequences", "system", "metadata"]:
if param in kwargs:
create_kwargs[param] = kwargs[param]
response = self.client.messages.create(**create_kwargs)
return response.content[0].text
def generate_structured(self, prompt: str, **kwargs) -> dict:
@@ -703,11 +901,23 @@ class AnthropicProvider(BaseProvider):
raise ProcessingError("Anthropic client not initialized.")
json_prompt = f"{prompt}\n\nReturn the response as valid JSON only."
response = self.client.messages.create(
model=kwargs.get("model", self.model),
max_tokens=kwargs.get("max_tokens", 4096),
messages=[{"role": "user", "content": json_prompt}],
)
# Anthropic requires max_tokens.
max_tokens = kwargs.get("max_tokens", 8192)
# Prepare arguments
create_kwargs = {
"model": kwargs.get("model", self.model),
"max_tokens": max_tokens,
"messages": [{"role": "user", "content": json_prompt}],
}
# Pass through other common parameters
for param in ["temperature", "top_p", "top_k", "stop_sequences", "system", "metadata"]:
if param in kwargs:
create_kwargs[param] = kwargs[param]
response = self.client.messages.create(**create_kwargs)
try:
return self._parse_json(response.content[0].text)
except Exception as e:
@@ -758,10 +968,25 @@ class OllamaProvider(BaseProvider):
"Ollama client not initialized. Make sure Ollama is running."
)
options = {"temperature": kwargs.get("temperature", 0.3)}
if "max_tokens" in kwargs:
options["num_predict"] = kwargs["max_tokens"]
if "num_ctx" in kwargs:
options["num_ctx"] = kwargs["num_ctx"]
elif "context_window" in kwargs:
options["num_ctx"] = kwargs["context_window"]
# Pass through other common options
for param in ["top_p", "top_k", "repeat_penalty", "seed"]:
if param in kwargs:
options[param] = kwargs[param]
response = self.client.generate(
model=kwargs.get("model", self.model),
prompt=prompt,
options={"temperature": kwargs.get("temperature", 0.3)},
options=options,
)
return response.get("response", "")
@@ -771,10 +996,26 @@ class OllamaProvider(BaseProvider):
raise ProcessingError("Ollama client not initialized.")
json_prompt = f"{prompt}\n\nReturn the response as valid JSON only."
options = {"temperature": kwargs.get("temperature", 0.3)}
if "max_tokens" in kwargs:
options["num_predict"] = kwargs["max_tokens"]
if "num_ctx" in kwargs:
options["num_ctx"] = kwargs["num_ctx"]
elif "context_window" in kwargs:
options["num_ctx"] = kwargs["context_window"]
# Pass through other common options
for param in ["top_p", "top_k", "repeat_penalty", "seed"]:
if param in kwargs:
options[param] = kwargs[param]
response = self.client.generate(
model=kwargs.get("model", self.model),
prompt=json_prompt,
options={"temperature": kwargs.get("temperature", 0.3)},
options=options,
)
try:
return self._parse_json(response.get("response", "{}"))
@@ -808,11 +1049,16 @@ class DeepSeekProvider(BaseProvider):
def generate(self, prompt: str, **kwargs) -> str:
if not self.client:
raise ProcessingError("DeepSeek client not initialized. Set DEEPSEEK_API_KEY or pass api_key.")
response = self.client.chat.completions.create(
model=kwargs.get("model", self.model),
messages=[{"role": "user", "content": prompt}],
temperature=kwargs.get("temperature", 0.3),
)
create_kwargs = {
"model": kwargs.get("model", self.model),
"messages": [{"role": "user", "content": prompt}],
"temperature": kwargs.get("temperature", 0.3),
}
if "max_tokens" in kwargs:
create_kwargs["max_tokens"] = kwargs["max_tokens"]
response = self.client.chat.completions.create(**create_kwargs)
return response.choices[0].message.content
def generate_structured(self, prompt: str, **kwargs) -> Union[dict, list]:
"""Generate structured output."""
@@ -879,11 +1125,27 @@ class HuggingFaceLLMProvider(BaseProvider):
raise ProcessingError("HuggingFace model not initialized.")
inputs = self.tokenizer.encode(prompt, return_tensors="pt").to(self.device)
# Use max_new_tokens if available, otherwise fallback to max_length with a safe default
generate_kwargs = {
"temperature": kwargs.get("temperature", 0.7),
"do_sample": True,
}
if "max_new_tokens" in kwargs:
generate_kwargs["max_new_tokens"] = kwargs["max_new_tokens"]
elif "max_tokens" in kwargs:
generate_kwargs["max_new_tokens"] = kwargs["max_tokens"]
# Support legacy max_length if explicitly provided
if "max_length" in kwargs:
generate_kwargs["max_length"] = kwargs["max_length"]
# Remove max_new_tokens if max_length is set to avoid conflict
generate_kwargs.pop("max_new_tokens", None)
outputs = self.model.generate(
inputs,
max_length=kwargs.get("max_length", 100),
temperature=kwargs.get("temperature", 0.7),
do_sample=True,
**generate_kwargs
)
generated_text = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
# Remove the original prompt from the response
@@ -1007,16 +1269,33 @@ class HuggingFaceModelLoader:
# This would need to be customized based on the model architecture
return model(text)
def extract_triplets(self, model, text: str) -> List[Dict]:
def extract_triplets(self, model, text: str, **kwargs) -> List[Dict]:
"""Extract triplets using loaded model."""
tokenizer = model["tokenizer"]
model_obj = model["model"]
device = model["device"]
# Use kwargs for max_length, default to 512 for input and 128 for output if not specified
max_input_length = kwargs.get("max_input_length", 512)
max_length = kwargs.get("max_length", 128)
# Allow max_new_tokens as well
generate_kwargs = {"max_length": max_length}
if "max_new_tokens" in kwargs:
generate_kwargs["max_new_tokens"] = kwargs["max_new_tokens"]
# If max_new_tokens is set, we might want to remove max_length or ensure they don't conflict
# For Seq2Seq, max_length usually refers to the total length of the target sequence
# Pass other generation args
for param in ["num_beams", "temperature", "top_p", "top_k", "do_sample"]:
if param in kwargs:
generate_kwargs[param] = kwargs[param]
inputs = tokenizer(
text, return_tensors="pt", truncation=True, max_length=512
text, return_tensors="pt", truncation=True, max_length=max_input_length
).to(device)
outputs = model_obj.generate(**inputs, max_length=128)
outputs = model_obj.generate(**inputs, **generate_kwargs)
decoded = tokenizer.decode(outputs[0], skip_special_tokens=True)
# Parse decoded output (format depends on model)
@@ -1024,28 +1303,88 @@ class HuggingFaceModelLoader:
return [{"triplet": decoded}]
def create_provider(name: str, **kwargs) -> BaseProvider:
"""Create provider - checks registry for custom providers."""
# Check registry first
custom_provider = provider_registry.get(name)
if custom_provider:
return custom_provider(**kwargs)
class ProviderPool:
"""Pool for reusing provider instances."""
def __init__(self):
self._providers: Dict[str, BaseProvider] = {}
self.logger = get_logger("provider_pool")
# Built-in providers
builtin = {
"openai": OpenAIProvider,
"gemini": GeminiProvider,
"groq": GroqProvider,
"anthropic": AnthropicProvider,
"ollama": OllamaProvider,
"huggingface_llm": HuggingFaceLLMProvider,
"deepseek": DeepSeekProvider,
}
def get(self, name: str, **kwargs) -> BaseProvider:
"""Get or create a provider instance."""
# Create a cache key from name and kwargs
# Filter out non-hashable items or volatile args if any
# For now, we assume kwargs are configuration options that should match
# Helper to make dict hashable
def make_hashable(value):
if isinstance(value, dict):
return tuple(sorted((k, make_hashable(v)) for k, v in value.items()))
elif isinstance(value, list):
return tuple(make_hashable(v) for v in value)
return value
provider_class = builtin.get(name.lower())
if not provider_class:
raise ValueError(
f"Unknown provider: {name}. Register custom provider or use built-in: {list(builtin.keys())}"
)
key_parts = [name]
for k, v in sorted(kwargs.items()):
# Skip some keys if they shouldn't affect pooling?
# For now, all init args matter for the instance identity.
key_parts.append((k, make_hashable(v)))
key = str(tuple(key_parts))
if key in self._providers:
return self._providers[key]
self.logger.debug(f"Creating new provider instance for {name}")
provider = self._create_provider(name, **kwargs)
self._providers[key] = provider
return provider
def _create_provider(self, name: str, **kwargs) -> BaseProvider:
"""Internal creation logic."""
# Check registry first
custom_provider = provider_registry.get(name)
if custom_provider:
return custom_provider(**kwargs)
return provider_class(**kwargs)
# Built-in providers
builtin = {
"openai": OpenAIProvider,
"gemini": GeminiProvider,
"groq": GroqProvider,
"anthropic": AnthropicProvider,
"ollama": OllamaProvider,
"huggingface_llm": HuggingFaceLLMProvider,
"deepseek": DeepSeekProvider,
}
provider_class = builtin.get(name.lower())
if not provider_class:
raise ValueError(
f"Unknown provider: {name}. Register custom provider or use built-in: {list(builtin.keys())}"
)
return provider_class(**kwargs)
def clear(self):
"""Clear the provider pool."""
self._providers.clear()
# Global provider pool
_provider_pool = ProviderPool()
def create_provider(name: str, use_pool: bool = True, **kwargs) -> BaseProvider:
"""
Create provider - checks registry for custom providers.
Args:
name: Provider name
use_pool: Whether to use the provider pool (default: True)
**kwargs: Provider arguments
"""
if use_pool:
return _provider_pool.get(name, **kwargs)
return _provider_pool._create_provider(name, **kwargs)
+128 -59
View File
@@ -201,10 +201,12 @@ class RelationExtractor:
)
try:
results = []
# Ensure lists are same length
min_len = min(len(text), len(entities))
results = [None] * min_len
total_relations_count = 0
processed_count = 0
# Update more frequently: every 1% or at least every 10 items, but always update for small datasets
if min_len <= 10:
update_interval = 1 # Update every item for small datasets
@@ -212,58 +214,101 @@ class RelationExtractor:
update_interval = max(1, min(10, min_len // 100))
# Initial progress update - ALWAYS show this
remaining = min_len
self.progress_tracker.update_progress(
tracking_id,
processed=0,
total=min_len,
message=f"Starting batch extraction... 0/{min_len} (remaining: {remaining})"
message=f"Starting batch extraction... 0/{min_len}"
)
for i in range(min_len):
doc_item = text[i]
ent_item = entities[i]
doc_text = ""
if isinstance(doc_item, dict) and "content" in doc_item:
doc_text = doc_item["content"]
elif isinstance(doc_item, str):
doc_text = doc_item
else:
doc_text = str(doc_item)
# Ensure ent_item is a list of entities
if not isinstance(ent_item, list):
ent_item = [] # Should not happen if entities is List[List[Entity]]
current_relations = self.extract_relations(doc_text, ent_item, **kwargs)
# Add provenance metadata
for rel in current_relations:
if rel.metadata is None:
rel.metadata = {}
rel.metadata["batch_index"] = i
if isinstance(doc_item, dict) and "id" in doc_item:
rel.metadata["document_id"] = doc_item["id"]
from .config import resolve_max_workers
max_workers = resolve_max_workers(
explicit=kwargs.get("max_workers"),
local_config=self.config,
methods=self.method,
)
results.append(current_relations)
total_relations_count += len(current_relations)
def process_item(i, doc_item, ent_item):
try:
doc_text = ""
if isinstance(doc_item, dict) and "content" in doc_item:
doc_text = doc_item["content"]
elif isinstance(doc_item, str):
doc_text = doc_item
else:
doc_text = str(doc_item)
# Ensure ent_item is a list of entities
if not isinstance(ent_item, list):
ent_item = [] # Should not happen if entities is List[List[Entity]]
current_relations = self.extract_relations(doc_text, ent_item, **kwargs)
# Add provenance metadata
for rel in current_relations:
if rel.metadata is None:
rel.metadata = {}
rel.metadata["batch_index"] = i
if isinstance(doc_item, dict) and "id" in doc_item:
rel.metadata["document_id"] = doc_item["id"]
return i, current_relations
except Exception as e:
self.logger.warning(f"Failed to process item {i}: {e}")
return i, []
if max_workers > 1:
import concurrent.futures
remaining = min_len - (i + 1)
# Update progress: always update for small datasets, or at intervals for large ones
should_update = (
(i + 1) % update_interval == 0 or
(i + 1) == min_len or
i == 0 or
min_len <= 10 # Always update for small datasets
)
if should_update:
self.progress_tracker.update_progress(
tracking_id,
processed=i + 1,
total=min_len,
message=f"Processing documents... {i + 1}/{min_len} (remaining: {remaining}) - Extracted {total_relations_count} relations so far"
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
# Submit tasks
future_to_idx = {
executor.submit(process_item, i, text[i], entities[i]): i
for i in range(min_len)
}
for future in concurrent.futures.as_completed(future_to_idx):
i, relations = future.result()
results[i] = relations
total_relations_count += len(relations)
processed_count += 1
should_update = (
processed_count % update_interval == 0 or
processed_count == min_len or
processed_count == 1 or
min_len <= 10
)
if should_update:
remaining = min_len - processed_count
self.progress_tracker.update_progress(
tracking_id,
processed=processed_count,
total=min_len,
message=f"Processing documents... {processed_count}/{min_len} (remaining: {remaining}) - Extracted {total_relations_count} relations so far"
)
else:
# Sequential processing
for i in range(min_len):
_, relations = process_item(i, text[i], entities[i])
results[i] = relations
total_relations_count += len(relations)
processed_count += 1
should_update = (
processed_count % update_interval == 0 or
processed_count == min_len or
processed_count == 1 or
min_len <= 10
)
if should_update:
remaining = min_len - processed_count
self.progress_tracker.update_progress(
tracking_id,
processed=processed_count,
total=min_len,
message=f"Processing documents... {processed_count}/{min_len} (remaining: {remaining}) - Extracted {total_relations_count} relations so far"
)
self.progress_tracker.stop_tracking(
tracking_id,
@@ -284,14 +329,19 @@ class RelationExtractor:
return []
def extract_relations(
self, text: str, entities: List[Entity], **options
) -> List[Relation]:
self,
text: Union[str, List[Dict[str, Any]], List[str]],
entities: Union[List[Entity], List[List[Entity]]],
pipeline_id: Optional[str] = None,
**options,
) -> Union[List[Relation], List[List[Relation]]]:
"""
Extract relations between entities.
Args:
text: Input text
entities: List of extracted entities
pipeline_id: Optional pipeline ID for progress tracking (batch mode)
**options: Extraction options:
- method: Override method (if not set in __init__)
- min_confidence: Minimum confidence threshold
@@ -300,7 +350,18 @@ class RelationExtractor:
Returns:
list: List of extracted relations
"""
from .methods import get_relation_method
if isinstance(text, list):
if entities is None:
entities_batch = [[] for _ in range(len(text))]
elif isinstance(entities, list) and (not entities):
entities_batch = [[] for _ in range(len(text))]
elif isinstance(entities, list) and all(isinstance(e, Entity) for e in entities):
entities_batch = [entities for _ in range(len(text))]
else:
entities_batch = entities
return self.extract(text, entities_batch, pipeline_id=pipeline_id, **options)
from .methods import get_relation_method, match_entity
tracking_id = self.progress_tracker.start_tracking(
module="semantic_extract",
@@ -358,14 +419,13 @@ class RelationExtractor:
method_options["model"] = all_options.get(
"llm_model", all_options.get("model")
)
# Pass api_key if provided (needed for all providers)
if "api_key" in all_options:
method_options["api_key"] = all_options["api_key"]
elif "api_key" not in method_options:
# Try to get from environment as fallback
# Ensure api_key is populated: check explicitly provided or fallback to env
current_key = method_options.get("api_key")
if not current_key:
# Not found or empty/None, try environment
import os
provider = method_options.get("provider", "openai")
env_key = f"{provider.upper()}_API_KEY"
provider_name = method_options.get("provider", "openai")
env_key = f"{provider_name.upper()}_API_KEY"
api_key = os.getenv(env_key)
if api_key:
method_options["api_key"] = api_key
@@ -379,6 +439,12 @@ class RelationExtractor:
if verbose_mode and method_name == "llm":
import sys
print(f" [RelationExtractor] Processing with {method_name}...", flush=True, file=sys.stdout)
print(f" [RelationExtractor Debug] method_options keys: {list(method_options.keys())}", flush=True, file=sys.stdout)
if "api_key" in method_options:
masked = method_options["api_key"][:4] + "..." if method_options["api_key"] else "None"
print(f" [RelationExtractor Debug] api_key present: {masked}", flush=True, file=sys.stdout)
else:
print(f" [RelationExtractor Debug] api_key NOT present", flush=True, file=sys.stdout)
relations = method_func(text, entities, **method_options)
@@ -421,6 +487,11 @@ class RelationExtractor:
except Exception as e:
self.logger.warning(f"Method {method_name} failed: {e}")
if verbose_mode:
import sys
print(f" [RelationExtractor] ERROR: Method {method_name} failed: {e}", flush=True, file=sys.stderr)
import traceback
traceback.print_exc(file=sys.stderr)
continue
# Use first successful method or combine
@@ -487,11 +558,9 @@ class RelationExtractor:
self, text: str, entities: List[Entity]
) -> List[Relation]:
"""Extract relations using pattern matching."""
from .methods import match_entity
relations = []
# Create entity lookup by text
entity_map = {e.text.lower(): e for e in entities}
# Check each relation pattern
for relation_type, patterns in self.relation_patterns.items():
for pattern in patterns:
@@ -499,8 +568,8 @@ class RelationExtractor:
subject_text = match.group("subject").strip()
object_text = match.group("object").strip()
subject_entity = entity_map.get(subject_text.lower())
object_entity = entity_map.get(object_text.lower())
subject_entity = match_entity(subject_text, entities)
object_entity = match_entity(object_text, entities)
if subject_entity and object_entity:
# Get context around the match
+75 -31
View File
@@ -148,8 +148,9 @@ class SemanticAnalyzer:
)
try:
results = []
results = [None] * len(text)
total_items = len(text)
processed_count = 0
# Determine update interval
if total_items <= 10:
@@ -165,40 +166,83 @@ class SemanticAnalyzer:
message=f"Starting batch analysis... 0/{total_items} (remaining: {total_items})"
)
for idx, item in enumerate(text):
# Prepare arguments for single item
doc_text = item["content"] if isinstance(item, dict) and "content" in item else str(item)
# Analyze
analysis = self.analyze_semantics(doc_text, **kwargs)
from .config import resolve_max_workers
max_workers = resolve_max_workers(
explicit=kwargs.get("max_workers"),
local_config=self.config,
)
# Add provenance metadata
analysis["batch_index"] = idx
if isinstance(item, dict) and "id" in item:
analysis["document_id"] = item["id"]
# Also inject into semantic roles if present
if "semantic_roles" in analysis:
for role in analysis["semantic_roles"]:
# role is a dict here because analyze_semantics converts it
if "metadata" not in role:
role["metadata"] = {}
role["metadata"]["batch_index"] = idx
if isinstance(item, dict) and "id" in item:
role["metadata"]["document_id"] = item["id"]
def process_item(idx, item):
try:
doc_text = item["content"] if isinstance(item, dict) and "content" in item else str(item)
analysis = self.analyze_semantics(doc_text, **kwargs)
results.append(analysis)
analysis["batch_index"] = idx
if isinstance(item, dict) and "id" in item:
analysis["document_id"] = item["id"]
# Update progress
if (idx + 1) % update_interval == 0 or (idx + 1) == total_items:
remaining = total_items - (idx + 1)
self.progress_tracker.update_progress(
tracking_id,
processed=idx + 1,
total=total_items,
message=f"Processing... {idx + 1}/{total_items} (remaining: {remaining})"
if "semantic_roles" in analysis:
for role in analysis["semantic_roles"]:
if "metadata" not in role:
role["metadata"] = {}
role["metadata"]["batch_index"] = idx
if isinstance(item, dict) and "id" in item:
role["metadata"]["document_id"] = item["id"]
return idx, analysis
except Exception as e:
self.logger.warning(f"Failed to analyze item {idx}: {e}")
return idx, {"error": str(e), "batch_index": idx}
if max_workers > 1:
import concurrent.futures
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
future_to_idx = {
executor.submit(process_item, idx, item): idx
for idx, item in enumerate(text)
}
for future in concurrent.futures.as_completed(future_to_idx):
idx, analysis = future.result()
results[idx] = analysis
processed_count += 1
should_update = (
processed_count % update_interval == 0
or processed_count == total_items
or processed_count == 1
or total_items <= 10
)
if should_update:
remaining = total_items - processed_count
self.progress_tracker.update_progress(
tracking_id,
processed=processed_count,
total=total_items,
message=f"Processing... {processed_count}/{total_items} (remaining: {remaining})"
)
else:
for idx, item in enumerate(text):
_, analysis = process_item(idx, item)
results[idx] = analysis
processed_count += 1
should_update = (
processed_count % update_interval == 0
or processed_count == total_items
or processed_count == 1
or total_items <= 10
)
if should_update:
remaining = total_items - processed_count
self.progress_tracker.update_progress(
tracking_id,
processed=processed_count,
total=total_items,
message=f"Processing... {processed_count}/{total_items} (remaining: {remaining})"
)
self.progress_tracker.stop_tracking(
tracking_id,
@@ -41,6 +41,7 @@ print(f"Extracted {len(entities)} entities and {len(relations)} relations")
All extractors support batch processing for high-throughput extraction. You can pass a list of strings or a list of dictionaries (with `content` and `id` keys).
**Features:**
- **Parallel Processing**: Multi-threaded extraction for high throughput (control via `max_workers`).
- **Progress Tracking**: Automatically shows a progress bar for large batches.
- **Provenance Metadata**: Each extracted item includes `batch_index` and `document_id` in its `metadata`.
@@ -52,9 +53,13 @@ documents = [
{"id": "doc_2", "content": "Microsoft Corporation was founded by Bill Gates."}
]
extractor = NERExtractor()
# Initialize with parallel processing enabled
extractor = NERExtractor(max_workers=4)
batch_results = extractor.extract(documents)
# OR override during extraction call
# batch_results = extractor.extract(documents, max_workers=8)
for i, doc_entities in enumerate(batch_results):
print(f"Document {i} entities:")
for entity in doc_entities:
@@ -120,9 +125,22 @@ entities = extractor.extract(
provider="openai",
model="gpt-4",
silent_fail=False, # Raise ProcessingError on failure (default)
max_text_length=4000 # Auto-chunking for long text
max_text_length=4000, # Auto-chunking for long text (default: 64k for major providers)
max_tokens=4096, # Explicitly control generation output length
temperature=0.0
)
print(f"LLM method: {len(entities)} entities")
# Groq extraction with long context support
# Groq defaults to 64k chunking limit for models like llama-3.3-70b
groq_extractor = NERExtractor(method="llm")
groq_entities = groq_extractor.extract(
text,
provider="groq",
model="llama-3.3-70b-versatile",
max_tokens=8000 # Passed directly to Groq API
)
print(f"Groq method: {len(groq_entities)} entities")
```
### Using NERExtractor Directly
@@ -221,6 +239,8 @@ relations = extractor.extract(
text,
entities=entities,
provider="openai",
model="gpt-4",
max_tokens=2048, # Increased output limit for many relations
silent_fail=True # Return empty list if extraction fails
)
```
@@ -284,7 +304,8 @@ triplets = extractor.extract_triplets(
text,
provider="openai",
model="gpt-4",
max_text_length=2000 # Force chunking for long text
max_text_length=64000, # Large default chunk size supported
max_tokens=4096 # Ensure enough tokens for all triplets
)
```
@@ -149,6 +149,9 @@ class SemanticNetworkExtractor:
self.config["ner_method"] = method
self.config["relation_method"] = method
self._ner_extractor = None
self._relation_extractor = None
def extract(
self,
text: Union[str, List[str], List[Dict[str, Any]]],
@@ -181,8 +184,9 @@ class SemanticNetworkExtractor:
)
try:
results = []
results = [None] * len(text)
total_items = len(text)
processed_count = 0
# Determine update interval
if total_items <= 10:
@@ -198,54 +202,113 @@ class SemanticNetworkExtractor:
message=f"Starting batch extraction... 0/{total_items} (remaining: {total_items})"
)
for idx, item in enumerate(text):
# Prepare arguments for single item
doc_text = item["content"] if isinstance(item, dict) and "content" in item else str(item)
doc_entities = None
if entities and isinstance(entities, list) and idx < len(entities):
doc_entities = entities[idx]
doc_relations = None
if relations and isinstance(relations, list) and idx < len(relations):
doc_relations = relations[idx]
from .config import resolve_max_workers
max_workers = resolve_max_workers(
explicit=kwargs.get("max_workers"),
local_config=self.config,
)
# Extract
network = self.extract_network(
doc_text,
entities=doc_entities,
relations=doc_relations,
**kwargs
)
# Add provenance metadata to nodes and edges
batch_meta = {"batch_index": idx}
if isinstance(item, dict) and "id" in item:
batch_meta["document_id"] = item["id"]
# Update network metadata
network.metadata.update(batch_meta)
# Update nodes metadata
for node in network.nodes:
node.metadata.update(batch_meta)
def process_item(idx, item, doc_entities, doc_relations):
try:
doc_text = item["content"] if isinstance(item, dict) and "content" in item else str(item)
# Update edges metadata
for edge in network.edges:
edge.metadata.update(batch_meta)
results.append(network)
# Update progress
if (idx + 1) % update_interval == 0 or (idx + 1) == total_items:
remaining = total_items - (idx + 1)
self.progress_tracker.update_progress(
tracking_id,
processed=idx + 1,
total=total_items,
message=f"Processing... {idx + 1}/{total_items} (remaining: {remaining})"
# Extract
network = self.extract_network(
doc_text,
entities=doc_entities,
relations=doc_relations,
**kwargs
)
# Add provenance metadata to nodes and edges
batch_meta = {"batch_index": idx}
if isinstance(item, dict) and "id" in item:
batch_meta["document_id"] = item["id"]
# Update network metadata
network.metadata.update(batch_meta)
# Update nodes metadata
for node in network.nodes:
node.metadata.update(batch_meta)
# Update edges metadata
for edge in network.edges:
edge.metadata.update(batch_meta)
return idx, network
except Exception as e:
self.logger.warning(f"Failed to process item {idx}: {e}")
verbose_mode = kwargs.get("verbose", False) or self.config.get("verbose", False)
if verbose_mode:
import sys
print(f" [SemanticNetworkExtractor] ERROR: Batch item {idx} failed: {e}", flush=True, file=sys.stderr)
return idx, None
if max_workers > 1:
import concurrent.futures
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
# Submit tasks
future_to_idx = {}
for idx, item in enumerate(text):
doc_entities = None
if entities and isinstance(entities, list) and idx < len(entities):
doc_entities = entities[idx]
doc_relations = None
if relations and isinstance(relations, list) and idx < len(relations):
doc_relations = relations[idx]
future = executor.submit(process_item, idx, item, doc_entities, doc_relations)
future_to_idx[future] = idx
for future in concurrent.futures.as_completed(future_to_idx):
idx, network = future.result()
if network:
results[idx] = network
processed_count += 1
# Update progress
if (processed_count) % update_interval == 0 or (processed_count) == total_items:
remaining = total_items - processed_count
self.progress_tracker.update_progress(
tracking_id,
processed=processed_count,
total=total_items,
message=f"Processing... {processed_count}/{total_items} (remaining: {remaining})"
)
else:
# Sequential processing
for idx, item in enumerate(text):
doc_entities = None
if entities and isinstance(entities, list) and idx < len(entities):
doc_entities = entities[idx]
doc_relations = None
if relations and isinstance(relations, list) and idx < len(relations):
doc_relations = relations[idx]
_, network = process_item(idx, item, doc_entities, doc_relations)
if network:
results[idx] = network
processed_count += 1
# Update progress
if (processed_count) % update_interval == 0 or (processed_count) == total_items:
remaining = total_items - processed_count
self.progress_tracker.update_progress(
tracking_id,
processed=processed_count,
total=total_items,
message=f"Processing... {processed_count}/{total_items} (remaining: {remaining})"
)
# Filter out None results if any failed
results = [r for r in results if r is not None]
self.progress_tracker.stop_tracking(
tracking_id,
status="completed",
@@ -265,11 +328,12 @@ class SemanticNetworkExtractor:
def extract_network(
self,
text: str,
entities: Optional[List[Entity]] = None,
relations: Optional[List[Relation]] = None,
text: Union[str, List[str], List[Dict[str, Any]]],
entities: Optional[Union[List[Entity], List[List[Entity]]]] = None,
relations: Optional[Union[List[Relation], List[List[Relation]]]] = None,
pipeline_id: Optional[str] = None,
**options,
) -> SemanticNetwork:
) -> Union[SemanticNetwork, List[SemanticNetwork]]:
"""
Extract semantic network from text.
@@ -282,6 +346,23 @@ class SemanticNetworkExtractor:
Returns:
SemanticNetwork: Extracted semantic network
"""
if isinstance(text, list):
entities_batch = entities
if entities is not None and isinstance(entities, list) and (not entities or all(isinstance(e, Entity) for e in entities)):
entities_batch = [entities for _ in range(len(text))] if entities else [[] for _ in range(len(text))]
relations_batch = relations
if relations is not None and isinstance(relations, list) and (not relations or all(isinstance(r, Relation) for r in relations)):
relations_batch = [relations for _ in range(len(text))] if relations else [[] for _ in range(len(text))]
return self.extract(
text,
entities=entities_batch,
relations=relations_batch,
pipeline_id=pipeline_id,
**options,
)
tracking_id = self.progress_tracker.start_tracking(
module="semantic_extract",
submodule="SemanticNetworkExtractor",
@@ -301,15 +382,16 @@ class SemanticNetworkExtractor:
# Pass method if specified
if "ner_method" in self.config:
ner_config["method"] = self.config["ner_method"]
ner = NERExtractor(
**ner_config,
**{
k: v
for k, v in self.config.items()
if k not in ["ner", "relation"]
},
)
entities = ner.extract_entities(text, **options)
if self._ner_extractor is None:
self._ner_extractor = NERExtractor(
**ner_config,
**{
k: v
for k, v in self.config.items()
if k not in ["ner", "relation"]
},
)
entities = self._ner_extractor.extract_entities(text, **options)
# Extract relations if not provided
if relations is None:
@@ -320,15 +402,16 @@ class SemanticNetworkExtractor:
# Pass method if specified
if "relation_method" in self.config:
rel_config["method"] = self.config["relation_method"]
rel_extractor = RelationExtractor(
**rel_config,
**{
k: v
for k, v in self.config.items()
if k not in ["ner", "relation"]
},
)
relations = rel_extractor.extract_relations(text, entities, **options)
if self._relation_extractor is None:
self._relation_extractor = RelationExtractor(
**rel_config,
**{
k: v
for k, v in self.config.items()
if k not in ["ner", "relation"]
},
)
relations = self._relation_extractor.extract_relations(text, entities, **options)
# Build network
total_steps = 2 # Create nodes, create edges
@@ -355,6 +438,12 @@ class SemanticNetworkExtractor:
self.progress_tracker.stop_tracking(
tracking_id, status="failed", message=str(e)
)
verbose_mode = options.get("verbose", False) or self.config.get("verbose", False)
if verbose_mode:
import sys
print(f" [SemanticNetworkExtractor] ERROR: Extraction failed: {e}", flush=True, file=sys.stderr)
import traceback
traceback.print_exc(file=sys.stderr)
raise
def _build_network(
+165 -54
View File
@@ -143,6 +143,13 @@ class TripletExtractor:
if not self.progress_tracker.enabled:
self.progress_tracker.enabled = True
if method is not None:
self.config["ner_method"] = method
self.config["relation_method"] = method
self._ner_extractor = None
self._relation_extractor = None
# Store parameters
self.triplet_types = triplet_types
self.include_temporal = include_temporal
@@ -191,9 +198,10 @@ class TripletExtractor:
)
try:
results = []
results = [None] * len(text)
total_items = len(text)
total_triplets_count = 0
processed_count = 0
# Determine update interval
if total_items <= 10:
@@ -206,50 +214,91 @@ class TripletExtractor:
tracking_id,
processed=0,
total=total_items,
message=f"Starting batch extraction... 0/{total_items} (remaining: {total_items})"
message=f"Starting batch extraction... 0/{total_items}"
)
for idx, item in enumerate(text):
# Prepare arguments for single item
doc_text = item["content"] if isinstance(item, dict) and "content" in item else str(item)
doc_entities = None
if entities and isinstance(entities, list) and idx < len(entities):
doc_entities = entities[idx]
doc_relations = None
if relations and isinstance(relations, list) and idx < len(relations):
doc_relations = relations[idx]
from .config import resolve_max_workers
max_workers = resolve_max_workers(
explicit=kwargs.get("max_workers"),
local_config=self.config,
methods=self.method,
)
# Extract
current_triplets = self.extract_triplets(
doc_text,
entities=doc_entities,
relations=doc_relations,
**kwargs
)
def process_item(idx, item):
try:
# Prepare arguments for single item
doc_text = item["content"] if isinstance(item, dict) and "content" in item else str(item)
doc_entities = None
if entities and isinstance(entities, list) and idx < len(entities):
doc_entities = entities[idx]
doc_relations = None
if relations and isinstance(relations, list) and idx < len(relations):
doc_relations = relations[idx]
# Add provenance metadata
for triplet in current_triplets:
if triplet.metadata is None:
triplet.metadata = {}
triplet.metadata["batch_index"] = idx
if isinstance(item, dict) and "id" in item:
triplet.metadata["document_id"] = item["id"]
results.append(current_triplets)
total_triplets_count += len(current_triplets)
# Update progress
if (idx + 1) % update_interval == 0 or (idx + 1) == total_items:
remaining = total_items - (idx + 1)
self.progress_tracker.update_progress(
tracking_id,
processed=idx + 1,
total=total_items,
message=f"Processing... {idx + 1}/{total_items} (remaining: {remaining}) - Extracted {total_triplets_count} triplets"
# Extract
current_triplets = self.extract_triplets(
doc_text,
entities=doc_entities,
relations=doc_relations,
**kwargs
)
# Add provenance metadata
for triplet in current_triplets:
if triplet.metadata is None:
triplet.metadata = {}
triplet.metadata["batch_index"] = idx
if isinstance(item, dict) and "id" in item:
triplet.metadata["document_id"] = item["id"]
return idx, current_triplets
except Exception as e:
self.logger.warning(f"Failed to process item {idx}: {e}")
return idx, []
if max_workers > 1:
import concurrent.futures
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
# Submit tasks
future_to_idx = {
executor.submit(process_item, idx, item): idx
for idx, item in enumerate(text)
}
for future in concurrent.futures.as_completed(future_to_idx):
idx, triplets = future.result()
results[idx] = triplets
total_triplets_count += len(triplets)
processed_count += 1
if processed_count % update_interval == 0 or processed_count == total_items:
remaining = total_items - processed_count
self.progress_tracker.update_progress(
tracking_id,
processed=processed_count,
total=total_items,
message=f"Processing... {processed_count}/{total_items} (remaining: {remaining}) - Extracted {total_triplets_count} triplets"
)
else:
# Sequential processing
for idx, item in enumerate(text):
_, triplets = process_item(idx, item)
results[idx] = triplets
total_triplets_count += len(triplets)
processed_count += 1
if processed_count % update_interval == 0 or processed_count == total_items:
remaining = total_items - processed_count
self.progress_tracker.update_progress(
tracking_id,
processed=processed_count,
total=total_items,
message=f"Processing... {processed_count}/{total_items} (remaining: {remaining}) - Extracted {total_triplets_count} triplets"
)
self.progress_tracker.stop_tracking(
tracking_id,
status="completed",
@@ -269,11 +318,12 @@ class TripletExtractor:
def extract_triplets(
self,
text: str,
entities: Optional[List[Entity]] = None,
relations: Optional[List[Relation]] = None,
text: Union[str, List[str], List[Dict[str, Any]]],
entities: Optional[Union[List[Entity], List[List[Entity]]]] = None,
relations: Optional[Union[List[Relation], List[List[Relation]]]] = None,
pipeline_id: Optional[str] = None,
**options,
) -> List[Triplet]:
) -> Union[List[Triplet], List[List[Triplet]]]:
"""
Extract RDF triplets from text.
@@ -281,11 +331,29 @@ class TripletExtractor:
text: Input text
entities: Pre-extracted entities (optional)
relations: Pre-extracted relations (optional)
pipeline_id: Optional pipeline ID for progress tracking (batch mode)
**options: Extraction options
Returns:
list: List of extracted triplets
"""
if isinstance(text, list):
entities_batch = entities
if entities is not None and isinstance(entities, list) and (not entities or all(isinstance(e, Entity) for e in entities)):
entities_batch = [entities for _ in range(len(text))] if entities else [[] for _ in range(len(text))]
relations_batch = relations
if relations is not None and isinstance(relations, list) and (not relations or all(isinstance(r, Relation) for r in relations)):
relations_batch = [relations for _ in range(len(text))] if relations else [[] for _ in range(len(text))]
return self.extract(
text,
entities=entities_batch,
relations=relations_batch,
pipeline_id=pipeline_id,
**options,
)
from .methods import get_triplet_method
tracking_id = self.progress_tracker.start_tracking(
@@ -303,16 +371,38 @@ class TripletExtractor:
self.progress_tracker.update_tracking(
tracking_id, message="Extracting entities..."
)
ner = NERExtractor(**self.config.get("ner", {}))
entities = ner.extract_entities(text)
if self._ner_extractor is None:
ner_config = self.config.get("ner", {})
if "ner_method" in self.config:
ner_config = {**ner_config, "method": self.config["ner_method"]}
self._ner_extractor = NERExtractor(
**ner_config,
**{
k: v
for k, v in self.config.items()
if k not in ["ner", "relation", "validator", "serializer", "quality"]
},
)
entities = self._ner_extractor.extract_entities(text)
# Extract relations if not provided
if relations is None:
self.progress_tracker.update_tracking(
tracking_id, message="Extracting relations..."
)
rel_extractor = RelationExtractor(**self.config.get("relation", {}))
relations = rel_extractor.extract_relations(text, entities)
if self._relation_extractor is None:
rel_config = self.config.get("relation", {})
if "relation_method" in self.config:
rel_config = {**rel_config, "method": self.config["relation_method"]}
self._relation_extractor = RelationExtractor(
**rel_config,
**{
k: v
for k, v in self.config.items()
if k not in ["ner", "relation", "validator", "serializer", "quality"]
},
)
relations = self._relation_extractor.extract_relations(text, entities)
# Use method-based extraction
methods = options.get("method", self.method)
@@ -371,18 +461,28 @@ class TripletExtractor:
method_options["model"] = all_options.get(
"llm_model", all_options.get("model")
)
# Pass api_key if provided (needed for all providers)
if "api_key" in all_options:
method_options["api_key"] = all_options["api_key"]
elif "api_key" not in method_options:
# Try to get from environment as fallback
# Ensure api_key is populated: check explicitly provided or fallback to env
current_key = method_options.get("api_key")
if not current_key:
# Not found or empty/None, try environment
import os
provider = method_options.get("provider", "openai")
env_key = f"{provider.upper()}_API_KEY"
provider_name = method_options.get("provider", "openai")
env_key = f"{provider_name.upper()}_API_KEY"
api_key = os.getenv(env_key)
if api_key:
method_options["api_key"] = api_key
# Print progress if verbose mode is enabled (only for LLM method to avoid spam)
verbose_mode = options.get("verbose", False) or self.config.get("verbose", False)
if verbose_mode and method_name == "llm":
import sys
print(f" [TripletExtractor] Processing with {method_name}...", flush=True, file=sys.stdout)
if "api_key" in method_options:
masked = method_options["api_key"][:4] + "..." if method_options["api_key"] else "None"
print(f" [TripletExtractor Debug] api_key present: {masked}", flush=True, file=sys.stdout)
else:
print(f" [TripletExtractor Debug] api_key NOT present", flush=True, file=sys.stdout)
triplets = method_func(
text,
entities=entities,
@@ -390,6 +490,11 @@ class TripletExtractor:
**method_options,
)
# Print result count if verbose (only for LLM method)
if verbose_mode and method_name == "llm" and len(triplets) > 0:
import sys
print(f" [TripletExtractor] Extracted {len(triplets)} triplets", flush=True, file=sys.stdout)
# Apply weighted scoring if triplet_types are provided
if triplet_types:
try:
@@ -425,6 +530,12 @@ class TripletExtractor:
except Exception as e:
self.logger.warning(f"Method {method_name} failed: {e}")
verbose_mode = options.get("verbose", False) or self.config.get("verbose", False)
if verbose_mode:
import sys
print(f" [TripletExtractor] ERROR: Method {method_name} failed: {e}", flush=True, file=sys.stderr)
import traceback
traceback.print_exc(file=sys.stderr)
continue
# Use first successful method or fallback to relation conversion
+25 -4
View File
@@ -34,6 +34,7 @@ License: MIT
"""
import inspect
import os
import sys
import threading
import time
@@ -46,6 +47,13 @@ from typing import Any, Callable, Dict, List, Optional, Tuple, Union
from .logging import get_logger
DISABLE_JUPYTER_PROGRESS = os.getenv("SEMANTICA_DISABLE_JUPYTER_PROGRESS", "").strip().lower() in (
"1",
"true",
"yes",
"on",
)
# Try to import IPython for Jupyter support
try:
from IPython import get_ipython
@@ -1017,6 +1025,7 @@ class ProgressTracker:
# Detect environment - will be checked dynamically
self.is_jupyter = self._detect_jupyter()
self.disable_jupyter_progress = DISABLE_JUPYTER_PROGRESS
# Create displays
self.displays: List[ProgressDisplay] = []
@@ -1024,7 +1033,7 @@ class ProgressTracker:
# Always try Jupyter first if available, fallback to console
if IPYTHON_AVAILABLE:
# Try to detect Jupyter - if available, use it
if self.is_jupyter:
if self.is_jupyter and not self.disable_jupyter_progress:
self.displays.append(JupyterProgressDisplay(use_emoji=use_emoji))
# Also add console as fallback for immediate feedback
self.displays.append(
@@ -1213,7 +1222,11 @@ class ProgressTracker:
if IPYTHON_AVAILABLE and not self.is_jupyter:
self.is_jupyter = self._detect_jupyter()
# If Jupyter is now detected and we don't have a Jupyter display, add it
if self.is_jupyter and not any(isinstance(d, JupyterProgressDisplay) for d in self.displays):
if (
self.is_jupyter
and not self.disable_jupyter_progress
and not any(isinstance(d, JupyterProgressDisplay) for d in self.displays)
):
# Insert Jupyter display at the beginning for priority
self.displays.insert(0, JupyterProgressDisplay(use_emoji=self.use_emoji))
@@ -1341,7 +1354,11 @@ class ProgressTracker:
if IPYTHON_AVAILABLE and not self.is_jupyter:
self.is_jupyter = self._detect_jupyter()
# If Jupyter is now detected and we don't have a Jupyter display, add it
if self.is_jupyter and not any(isinstance(d, JupyterProgressDisplay) for d in self.displays):
if (
self.is_jupyter
and not self.disable_jupyter_progress
and not any(isinstance(d, JupyterProgressDisplay) for d in self.displays)
):
# Insert Jupyter display at the beginning for priority
self.displays.insert(0, JupyterProgressDisplay(use_emoji=self.use_emoji))
@@ -1526,7 +1543,11 @@ def get_progress_tracker() -> ProgressTracker:
if IPYTHON_AVAILABLE and not _global_tracker.is_jupyter:
_global_tracker.is_jupyter = _global_tracker._detect_jupyter()
# If Jupyter is now detected and we don't have a Jupyter display, add it
if _global_tracker.is_jupyter and not any(isinstance(d, JupyterProgressDisplay) for d in _global_tracker.displays):
if (
_global_tracker.is_jupyter
and not _global_tracker.disable_jupyter_progress
and not any(isinstance(d, JupyterProgressDisplay) for d in _global_tracker.displays)
):
# Insert Jupyter display at the beginning for priority
_global_tracker.displays.insert(0, JupyterProgressDisplay(use_emoji=_global_tracker.use_emoji))
+138 -1
View File
@@ -38,6 +38,7 @@ License: MIT
"""
from typing import Any, Dict, List, Optional, Tuple, Union
import concurrent.futures
import numpy as np
@@ -61,7 +62,7 @@ class VectorStore:
SUPPORTED_BACKENDS = {"faiss", "weaviate", "qdrant", "milvus", "inmemory"}
def __init__(self, backend="faiss", config=None, **kwargs):
def __init__(self, backend="faiss", config=None, max_workers: int = 6, **kwargs):
"""Initialize vector store."""
if backend.lower() not in self.SUPPORTED_BACKENDS:
raise ValueError(
@@ -72,6 +73,7 @@ class VectorStore:
self.logger = get_logger("vector_store")
self.config = config or {}
self.config.update(kwargs)
self.max_workers = max_workers
self.progress_tracker = get_progress_tracker()
# Ensure progress tracker is enabled
if not self.progress_tracker.enabled:
@@ -127,6 +129,141 @@ class VectorStore:
self.logger.warning("Using random fallback embedding")
return np.random.rand(self.dimension).astype(np.float32)
def embed_batch(self, texts: List[str]) -> List[np.ndarray]:
"""
Generate embeddings for a list of texts using the internal embedder.
Args:
texts: List of texts to embed
Returns:
List of numpy arrays
"""
if self.embedder:
try:
# generate_embeddings handles list input
embeddings = self.embedder.generate_embeddings(texts)
# Ensure it returns a list of arrays (it returns 2D array or list)
if isinstance(embeddings, np.ndarray):
return list(embeddings)
return embeddings
except Exception as e:
self.logger.warning(f"Batch embedding generation failed: {e}")
# Fallback
self.logger.warning("Using random fallback embeddings for batch")
return [np.random.rand(self.dimension).astype(np.float32) for _ in texts]
def add_documents(
self,
documents: List[str],
metadata: Optional[List[Dict[str, Any]]] = None,
batch_size: int = 32,
parallel: bool = True,
**options,
) -> List[str]:
"""
Add multiple documents to the store with parallel embedding generation.
Args:
documents: List of document texts
metadata: List of metadata dictionaries
batch_size: Number of documents to process in one batch
parallel: Whether to use parallel processing for embeddings
**options: Additional options
Returns:
List[str]: Vector IDs
"""
if not documents:
return []
num_docs = len(documents)
metadata = metadata or [{} for _ in range(num_docs)]
if len(metadata) != num_docs:
raise ValueError("Metadata list length must match documents length")
all_vectors = [None] * num_docs
# Helper for processing a batch
def process_batch(start_idx: int, end_idx: int):
batch_texts = documents[start_idx:end_idx]
batch_embeddings = self.embed_batch(batch_texts)
return start_idx, batch_embeddings
# Calculate batches
batches = []
for i in range(0, num_docs, batch_size):
batches.append((i, min(i + batch_size, num_docs)))
tracking_id = self.progress_tracker.start_tracking(
module="vector_store",
submodule="VectorStore",
message=f"Processing {num_docs} documents (parallel={parallel})",
)
try:
if parallel and self.max_workers > 1:
self.progress_tracker.update_tracking(
tracking_id, message=f"Embedding with {self.max_workers} workers..."
)
with concurrent.futures.ThreadPoolExecutor(max_workers=self.max_workers) as executor:
futures = [
executor.submit(process_batch, start, end)
for start, end in batches
]
completed = 0
for future in concurrent.futures.as_completed(futures):
start_idx, embeddings = future.result()
# Place results in correct order
for i, emb in enumerate(embeddings):
all_vectors[start_idx + i] = emb
completed += 1
if completed % 5 == 0: # Update progress periodically
self.progress_tracker.update_tracking(
tracking_id,
message=f"Embedded batch {completed}/{len(batches)}"
)
else:
# Sequential processing
self.progress_tracker.update_tracking(
tracking_id, message="Embedding sequentially..."
)
for i, (start, end) in enumerate(batches):
_, embeddings = process_batch(start, end)
for j, emb in enumerate(embeddings):
all_vectors[start + j] = emb
if i % 5 == 0:
self.progress_tracker.update_tracking(
tracking_id,
message=f"Embedded batch {i+1}/{len(batches)}"
)
# Verify all embeddings generated
if any(v is None for v in all_vectors):
raise ProcessingError("Failed to generate all embeddings")
# Store all vectors in one go
self.progress_tracker.update_tracking(tracking_id, message="Storing vectors...")
vector_ids = self.store_vectors(all_vectors, metadata=metadata, **options)
self.progress_tracker.stop_tracking(
tracking_id,
status="completed",
message=f"Added {len(vector_ids)} documents",
)
return vector_ids
except Exception as e:
self.progress_tracker.stop_tracking(
tracking_id, status="failed", message=str(e)
)
raise
def store(
self,
vectors: List[np.ndarray],
+214
View File
@@ -0,0 +1,214 @@
import os
import shutil
import tempfile
import pytest
from pathlib import Path
from semantica.ingest import OntologyIngestor, ingest, ingest_ontology, OntologyData
class TestOntologyIngestor:
@pytest.fixture
def sample_ttl_content(self):
return """
@prefix : <http://example.org/ontology/> .
@prefix owl: <http://www.w3.org/2002/07/owl#> .
@prefix rdf: <http://www.w3.org/1999/02/22-rdf-syntax-ns#> .
@prefix rdfs: <http://www.w3.org/2000/01/rdf-schema#> .
@prefix xsd: <http://www.w3.org/2001/XMLSchema#> .
<http://example.org/ontology/> rdf:type owl:Ontology ;
rdfs:label "Test Ontology" .
:Person rdf:type owl:Class ;
rdfs:label "Person" .
:hasName rdf:type owl:DatatypeProperty ;
rdfs:domain :Person ;
rdfs:range xsd:string .
"""
def test_ingest_single_file(self, sample_ttl_content):
with tempfile.NamedTemporaryFile(delete=False, suffix=".ttl", mode="w") as tmp:
tmp.write(sample_ttl_content)
tmp_path = tmp.name
try:
ingestor = OntologyIngestor()
result = ingestor.ingest_ontology(tmp_path)
assert isinstance(result, OntologyData)
assert result.data["name"] == "Test Ontology" or result.data["name"] == os.path.basename(tmp_path)
assert any(cls["name"] == "Person" for cls in result.data["classes"])
assert any(prop["name"] == "hasName" for prop in result.data["properties"])
assert result.metadata["format"] == "ttl" or result.metadata["format"] == "turtle"
finally:
if os.path.exists(tmp_path):
os.remove(tmp_path)
def test_ingest_directory(self, sample_ttl_content):
with tempfile.TemporaryDirectory() as tmp_dir:
# Create two ontology files
file1 = os.path.join(tmp_dir, "ont1.ttl")
file2 = os.path.join(tmp_dir, "ont2.rdf")
with open(file1, "w") as f:
f.write(sample_ttl_content)
# Simple RDF/XML content for the second file
rdf_content = """
<rdf:RDF xmlns:rdf="http://www.w3.org/1999/02/22-rdf-syntax-ns#"
xmlns:owl="http://www.w3.org/2002/07/owl#">
<owl:Ontology rdf:about="http://example.org/ont2"/>
<owl:Class rdf:about="http://example.org/ont2/Animal"/>
</rdf:RDF>
"""
with open(file2, "w") as f:
f.write(rdf_content)
ingestor = OntologyIngestor()
results = ingestor.ingest_directory(tmp_dir)
assert len(results) == 2
assert all(isinstance(r, OntologyData) for r in results)
# Verify results contain expected classes
classes = [cls["name"] for res in results for cls in res.data["classes"]]
assert "Person" in classes
assert "Animal" in classes
def test_unified_ingest_function(self, sample_ttl_content):
with tempfile.NamedTemporaryFile(delete=False, suffix=".ttl", mode="w") as tmp:
tmp.write(sample_ttl_content)
tmp_path = tmp.name
try:
# Test auto-detection via unified ingest
result = ingest(tmp_path)
assert "ontology" in result
assert isinstance(result["ontology"], OntologyData)
assert len(result["ontology"].data["classes"]) > 0
# Test explicit source type
result_explicit = ingest(tmp_path, source_type="ontology")
assert "ontology" in result_explicit
assert result_explicit["ontology"].metadata["source_path"] == tmp_path
finally:
if os.path.exists(tmp_path):
os.remove(tmp_path)
def test_convenience_function(self, sample_ttl_content):
with tempfile.NamedTemporaryFile(delete=False, suffix=".n3", mode="w") as tmp:
tmp.write(sample_ttl_content)
tmp_path = tmp.name
try:
result = ingest_ontology(tmp_path)
assert isinstance(result, OntologyData)
assert len(result.data["classes"]) > 0
finally:
if os.path.exists(tmp_path):
os.remove(tmp_path)
def test_ingest_formats(self):
"""Test ingestion of all supported formats."""
ingestor = OntologyIngestor()
# 1. JSON-LD
jsonld_content = """
{
"@context": {
"owl": "http://www.w3.org/2002/07/owl#",
"rdf": "http://www.w3.org/1999/02/22-rdf-syntax-ns#",
"rdfs": "http://www.w3.org/2000/01/rdf-schema#"
},
"@id": "http://example.org/jsonld",
"@type": "owl:Ontology",
"rdfs:label": "JSON-LD Ontology",
"owl:versionInfo": "1.0"
}
"""
with tempfile.NamedTemporaryFile(delete=False, suffix=".jsonld", mode="w") as tmp:
tmp.write(jsonld_content)
tmp_path = tmp.name
try:
result = ingestor.ingest_ontology(tmp_path)
assert result.data["name"] == "JSON-LD Ontology"
assert result.metadata["format"] == "json-ld"
finally:
if os.path.exists(tmp_path):
os.remove(tmp_path)
# 2. N-Triples
nt_content = '<http://example.org/nt/Class> <http://www.w3.org/1999/02/22-rdf-syntax-ns#type> <http://www.w3.org/2002/07/owl#Class> .\n'
with tempfile.NamedTemporaryFile(delete=False, suffix=".nt", mode="w") as tmp:
tmp.write(nt_content)
tmp_path = tmp.name
try:
result = ingestor.ingest_ontology(tmp_path)
# N-Triples often doesn't have ontology metadata, so name might default to basename
assert result.data["name"] == os.path.basename(tmp_path)
assert any(cls["uri"] == "http://example.org/nt/Class" for cls in result.data["classes"])
assert result.metadata["format"] == "nt"
finally:
if os.path.exists(tmp_path):
os.remove(tmp_path)
# 3. Notation3
n3_content = """
@prefix : <http://example.org/n3/> .
@prefix owl: <http://www.w3.org/2002/07/owl#> .
:N3Class a owl:Class .
"""
with tempfile.NamedTemporaryFile(delete=False, suffix=".n3", mode="w") as tmp:
tmp.write(n3_content)
tmp_path = tmp.name
try:
result = ingestor.ingest_ontology(tmp_path)
assert any(cls["uri"] == "http://example.org/n3/N3Class" for cls in result.data["classes"])
# format might be 'n3' or 'turtle' depending on rdflib detection as they are similar
assert result.metadata["format"] in ["n3", "turtle"]
finally:
if os.path.exists(tmp_path):
os.remove(tmp_path)
# 4. RDF/XML (.owl)
owl_content = """
<rdf:RDF xmlns:rdf="http://www.w3.org/1999/02/22-rdf-syntax-ns#"
xmlns:owl="http://www.w3.org/2002/07/owl#"
xmlns:rdfs="http://www.w3.org/2000/01/rdf-schema#">
<owl:Ontology rdf:about="http://example.org/owl"/>
<owl:Class rdf:about="http://example.org/owl/OwlClass">
<rdfs:label>OwlClass</rdfs:label>
</owl:Class>
</rdf:RDF>
"""
with tempfile.NamedTemporaryFile(delete=False, suffix=".owl", mode="w") as tmp:
tmp.write(owl_content)
tmp_path = tmp.name
try:
result = ingestor.ingest_ontology(tmp_path)
assert any(cls["name"] == "OwlClass" for cls in result.data["classes"])
assert result.metadata["format"] in ["xml", "rdf", "owl"]
finally:
if os.path.exists(tmp_path):
os.remove(tmp_path)
def test_error_handling(self):
ingestor = OntologyIngestor()
with pytest.raises(Exception): # Specific exception type depends on implementation, likely ValidationError or FileNotFoundError
ingestor.ingest_ontology("non_existent_file.ttl")
def test_invalid_content(self):
with tempfile.NamedTemporaryFile(delete=False, suffix=".ttl", mode="w") as tmp:
tmp.write("This is not valid turtle content")
tmp_path = tmp.name
try:
ingestor = OntologyIngestor()
# Depending on implementation, this might raise an exception or return partial/empty result with error in metadata
# Given current implementation uses g.parse(), it likely raises an exception which is caught or propagated
# If propagated:
with pytest.raises(Exception):
ingestor.ingest_ontology(tmp_path)
finally:
if os.path.exists(tmp_path):
os.remove(tmp_path)
+98
View File
@@ -0,0 +1,98 @@
import os
import unittest
import time
from semantica.semantic_extract import NERExtractor, RelationExtractor
# Use environment variable for API key
_GROQ_KEY = os.getenv("GROQ_API_KEY") or os.getenv("GROQ_TEST_API_KEY")
@unittest.skipUnless(_GROQ_KEY, "Groq key not set; skipping live integration test")
class TestGroqRelationsIntegration(unittest.TestCase):
def setUp(self):
self.api_key = _GROQ_KEY
self.model = "llama-3.1-8b-instant"
# Short, unambiguous finance snippet
self.text_short = (
"Apple reported revenue of $4.4 billion in Q1 2024 and provided guidance for FY 2025."
)
# Longer text to exercise chunking and ensure no hang
self.text_long = (
"Apple reported revenue of $4.4 billion in Q1 2024. "
"The company also reported growth of 12% year-over-year and provided guidance for FY 2025. "
"Microsoft reported revenue of $6.1 billion in Q2 2024 and expects sequential growth. "
"NVIDIA reported record revenue in 2024 Q1 and guided for higher revenue in Q2 2024. "
) * 20 # expand length
def _extract_entities(self, text):
ner = NERExtractor(
method="llm",
provider="groq",
llm_model=self.model,
api_key=self.api_key,
temperature=0.0,
)
entities = ner.extract_entities(text, entity_types=["ORGANIZATION", "MONEY", "DATE", "EVENT", "PERCENT"])
self.assertIsInstance(entities, list)
return entities
def test_relations_short_text(self):
entities = self._extract_entities(self.text_short)
self.assertGreater(len(entities), 0, "NER should extract entities for short text")
relation_extractor = RelationExtractor(
method="llm",
relation_types=[
"HAS_REVENUE",
"HAS_GROWTH",
"PROVIDES_GUIDANCE",
"IN_QUARTER",
"FOR_PERIOD",
"RELATED_TO",
],
provider="groq",
llm_model=self.model,
api_key=self.api_key,
temperature=0.0,
verbose=True,
)
start = time.time()
relations = relation_extractor.extract_relations(text=self.text_short, entities=entities)
elapsed = time.time() - start
self.assertIsInstance(relations, list)
# Ensure call completes reasonably fast (network dependent; allow generous bound)
self.assertLess(elapsed, 60, f"Extraction took too long: {elapsed:.2f}s")
# Do not strictly assert >0 as model output may vary, but log for diagnostics
if relations:
sample = relations[0]
self.assertTrue(hasattr(sample, "subject") and hasattr(sample, "predicate") and hasattr(sample, "object"))
def test_relations_long_text_chunking(self):
entities = self._extract_entities(self.text_long)
self.assertGreater(len(entities), 0, "NER should extract entities for long text")
relation_extractor = RelationExtractor(
method="llm",
relation_types=["RELATED_TO", "HAS_REVENUE", "IN_QUARTER"],
provider="groq",
llm_model=self.model,
api_key=self.api_key,
temperature=0.0,
verbose=True,
)
start = time.time()
relations = relation_extractor.extract_relations(text=self.text_long, entities=entities)
elapsed = time.time() - start
self.assertIsInstance(relations, list)
# Ensure completion (chunked path) and no hang
self.assertLess(elapsed, 120, f"Chunked extraction took too long: {elapsed:.2f}s")
if relations:
for r in relations[:3]:
self.assertTrue(hasattr(r, "subject") and hasattr(r, "predicate") and hasattr(r, "object"))
if __name__ == "__main__":
unittest.main()
+232
View File
@@ -0,0 +1,232 @@
import unittest
from unittest.mock import MagicMock, patch
from semantica.kg.graph_builder import GraphBuilder
class TestGraphBuilderExternal(unittest.TestCase):
def setUp(self):
self.mock_tracker_patcher = patch("semantica.utils.progress_tracker.get_progress_tracker")
self.mock_get_tracker = self.mock_tracker_patcher.start()
self.mock_tracker = MagicMock()
self.mock_get_tracker.return_value = self.mock_tracker
self.mock_resolver_patcher = patch("semantica.kg.entity_resolver.EntityResolver")
self.mock_resolver_cls = self.mock_resolver_patcher.start()
self.mock_conflict_patcher = patch("semantica.conflicts.conflict_detector.ConflictDetector")
self.mock_conflict_cls = self.mock_conflict_patcher.start()
def tearDown(self):
self.mock_tracker_patcher.stop()
self.mock_resolver_patcher.stop()
self.mock_conflict_patcher.stop()
def test_single_source_dict_with_source_id_target_id(self):
builder = GraphBuilder(merge_entities=False, resolve_conflicts=False)
entities = [
{"id": "drug:1", "name": "Aspirin", "type": "Drug"},
{"id": "disease:1", "name": "Myocardial infarction", "type": "Disease"},
]
relationships = [
{"source_id": "drug:1", "target_id": "disease:1", "type": "TREATS"},
]
source = {"entities": entities, "relationships": relationships}
kg = builder.build(source)
self.assertEqual(len(kg["entities"]), 2)
self.assertEqual(len(kg["relationships"]), 1)
rel = kg["relationships"][0]
self.assertEqual(rel.get("source"), "drug:1")
self.assertEqual(rel.get("target"), "disease:1")
self.assertEqual(kg["metadata"]["num_relationships"], 1)
def test_sources_list_merge_with_external_relationships(self):
builder = GraphBuilder(merge_entities=False, resolve_conflicts=False)
source1 = {
"entities": [{"id": "1", "name": "A"}],
"relationships": [{"source_id": "1", "target_id": "2", "type": "REL_1"}],
}
source2 = {
"entities": [{"id": "2", "name": "B"}],
"relationships": [{"source_id": "2", "target_id": "1", "type": "REL_2"}],
}
kg = builder.build([source1, source2])
self.assertEqual(len(kg["entities"]), 2)
self.assertEqual(len(kg["relationships"]), 2)
sources = {r["source"] for r in kg["relationships"]}
targets = {r["target"] for r in kg["relationships"]}
self.assertEqual(sources, {"1", "2"})
self.assertEqual(targets, {"1", "2"})
def test_build_with_explicit_relationships_argument_external_ids(self):
builder = GraphBuilder(merge_entities=False, resolve_conflicts=False)
entities = [
{"id": "1", "name": "A"},
{"id": "2", "name": "B"},
]
relationships = [
{"source_id": "1", "target_id": "2", "type": "REL"},
]
kg = builder.build(entities, relationships=relationships)
self.assertEqual(len(kg["entities"]), 2)
self.assertEqual(len(kg["relationships"]), 1)
rel = kg["relationships"][0]
self.assertEqual(rel.get("source"), "1")
self.assertEqual(rel.get("target"), "2")
def test_build_single_source_external_graph(self):
builder = GraphBuilder(merge_entities=False, resolve_conflicts=False)
source = {
"entities": [{"id": "1", "name": "A"}],
"relationships": [{"source_id": "1", "target_id": "1", "type": "SELF"}],
}
kg = builder.build_single_source(source)
self.assertEqual(len(kg["entities"]), 1)
self.assertEqual(len(kg["relationships"]), 1)
rel = kg["relationships"][0]
self.assertEqual(rel.get("source"), "1")
self.assertEqual(rel.get("target"), "1")
def test_relationship_key_variants_normalized(self):
builder = GraphBuilder(merge_entities=False, resolve_conflicts=False)
entities = [
{"id": "1", "name": "A"},
{"id": "2", "name": "B"},
{"id": "3", "name": "C"},
{"id": "4", "name": "D"},
]
relationships = [
{"source_id": "1", "target_id": "2", "type": "R1"},
{"source": "2", "target": "3", "type": "R2"},
{"subject": "3", "object": "4", "type": "R3"},
]
kg = builder.build({"entities": entities, "relationships": relationships})
self.assertEqual(len(kg["relationships"]), 3)
ids = {(r["source"], r["target"]) for r in kg["relationships"]}
self.assertIn(("1", "2"), ids)
self.assertIn(("2", "3"), ids)
self.assertIn(("3", "4"), ids)
def test_warning_when_all_relationships_dropped(self):
builder = GraphBuilder(merge_entities=False, resolve_conflicts=False)
source = {
"entities": [],
"relationships": [{"foo": "x"}, {"bar": "y"}],
}
with patch.object(builder.logger, "warning") as mock_warning:
kg = builder.build(source)
self.assertEqual(len(kg["relationships"]), 0)
mock_warning.assert_called()
args, _ = mock_warning.call_args
self.assertIn("All relationships were dropped", args[0])
def test_no_warning_when_some_relationships_kept(self):
builder = GraphBuilder(merge_entities=False, resolve_conflicts=False)
source = {
"entities": [{"id": "1"}, {"id": "2"}],
"relationships": [
{"source_id": "1", "target_id": "2", "type": "REL"},
{"foo": "x"},
],
}
with patch.object(builder.logger, "warning") as mock_warning:
kg = builder.build(source)
self.assertEqual(len(kg["relationships"]), 2)
mock_warning.assert_not_called()
def test_issue_208_minimal_reproduction_shape(self):
builder = GraphBuilder(
merge_entities=False,
entity_resolution_strategy="none",
resolve_conflicts=False,
)
entities = [
{"id": "e1", "name": "Entity 1"},
{"id": "e2", "name": "Entity 2"},
{"id": "e3", "name": "Entity 3"},
]
relationships = [
{"source_id": "e1", "target_id": "e2", "type": "REL_1"},
{"source_id": "e2", "target_id": "e3", "type": "REL_2"},
]
entity_ids = {e["id"] for e in entities}
for r in relationships:
self.assertIn(r["source_id"], entity_ids)
self.assertIn(r["target_id"], entity_ids)
kg = builder.build(
sources=[{"entities": entities, "relationships": relationships}],
merge_entities=False,
)
self.assertEqual(len(kg["entities"]), 3)
self.assertEqual(len(kg["relationships"]), 2)
pairs = {(r["source"], r["target"]) for r in kg["relationships"]}
self.assertIn(("e1", "e2"), pairs)
self.assertIn(("e2", "e3"), pairs)
def test_issue_206_earnings_call_shape(self):
builder = GraphBuilder(
merge_entities=False,
entity_resolution_strategy="none",
resolve_conflicts=False,
)
entities = [
{
"id": "entity_446_MDA Space Ltd.",
"name": "MDA Space Ltd.",
"type": "ORGANIZATION",
},
{
"id": "entity_500_$409.8 million",
"name": "$409.8 million",
"type": "MONEY",
},
]
relationships = [
{
"id": None,
"source_id": "MDA Space Ltd.",
"target_id": "$409.8 million",
"type": "HAS_REVENUE",
"confidence": 0.975,
"metadata": {},
}
]
kg = builder.build(
sources=[{"entities": entities, "relationships": relationships}],
merge_entities=False,
)
self.assertEqual(len(kg["entities"]), 2)
self.assertEqual(len(kg["relationships"]), 1)
rel = kg["relationships"][0]
self.assertEqual(rel.get("source"), "MDA Space Ltd.")
self.assertEqual(rel.get("target"), "$409.8 million")
+68
View File
@@ -91,6 +91,30 @@ class TestGraphBuilder(unittest.TestCase):
graph2 = builder.build(source_list)
self.assertEqual(len(graph2["entities"]), 2)
def test_build_with_external_relationship_ids(self):
builder = GraphBuilder(merge_entities=False, resolve_conflicts=False)
entities = [
{"id": "1", "name": "A"},
{"id": "2", "name": "B"},
]
relationships = [
{"source_id": "1", "target_id": "2", "type": "rel"},
]
source = {
"entities": entities,
"relationships": relationships,
}
graph = builder.build(source)
self.assertEqual(len(graph["entities"]), 2)
self.assertEqual(len(graph["relationships"]), 1)
rel = graph["relationships"][0]
self.assertEqual(rel.get("source"), "1")
self.assertEqual(rel.get("target"), "2")
def test_build_with_conflict_resolution(self):
"""Test building with conflict resolution enabled"""
builder = GraphBuilder(resolve_conflicts=True)
@@ -106,6 +130,50 @@ class TestGraphBuilder(unittest.TestCase):
self.mock_conflict_cls.return_value.detect_conflicts.assert_called_once()
self.mock_conflict_cls.return_value.resolve_conflicts.assert_called_once()
def test_build_single_source(self):
builder = GraphBuilder(merge_entities=False, resolve_conflicts=False)
source = {
"entities": [{"id": "1", "name": "A"}],
"relationships": [{"source_id": "1", "target_id": "1", "type": "self"}],
}
graph = builder.build_single_source(source)
self.assertEqual(len(graph["entities"]), 1)
self.assertEqual(len(graph["relationships"]), 1)
def test_build_with_explicit_relationships_argument(self):
builder = GraphBuilder(merge_entities=False, resolve_conflicts=False)
entities = [
{"id": "1", "name": "A"},
{"id": "2", "name": "B"},
]
relationships = [
{"source_id": "1", "target_id": "2", "type": "rel"},
]
graph = builder.build(entities, relationships=relationships)
self.assertEqual(len(graph["entities"]), 2)
self.assertEqual(len(graph["relationships"]), 1)
rel = graph["relationships"][0]
self.assertEqual(rel.get("source"), "1")
self.assertEqual(rel.get("target"), "2")
def test_build_warns_when_all_relationships_dropped(self):
builder = GraphBuilder(merge_entities=False, resolve_conflicts=False)
source = {
"entities": [],
"relationships": [{"foo": "x"}, {"bar": "y"}],
}
with patch.object(builder.logger, "warning") as mock_warning:
graph = builder.build(source)
self.assertEqual(len(graph["relationships"]), 0)
mock_warning.assert_called()
args, _ = mock_warning.call_args
self.assertIn("All relationships were dropped", args[0])
class TestGraphAnalyzer(unittest.TestCase):
def setUp(self):
self.mock_tracker_patcher = patch("semantica.kg.graph_analyzer.get_progress_tracker")
+20
View File
@@ -107,5 +107,25 @@ class TestPipelineModule(unittest.TestCase):
self.assertEqual(execution_order, ["A", "B", "C"])
def test_imports_no_circular_dependencies(self):
import semantica
_ = semantica.pipeline
from semantica.pipeline import PipelineBuilder, PipelineValidator
from semantica.deduplication import DuplicateDetector
builder = PipelineBuilder()
builder.add_step("step1", "dummy")
pipeline = builder.build("import_test_pipeline")
validator = PipelineValidator()
result = validator.validate_pipeline(pipeline)
self.assertTrue(result.valid)
detector = DuplicateDetector()
self.assertIsNotNone(detector)
if __name__ == "__main__":
unittest.main()
+100
View File
@@ -0,0 +1,100 @@
import unittest
from unittest.mock import MagicMock, patch
from semantica.semantic_extract.methods import extract_relations_llm, extract_entities_llm, extract_triplets_llm
from semantica.semantic_extract.ner_extractor import Entity
class TestMaxTokensPropagation(unittest.TestCase):
@patch("semantica.semantic_extract.methods.create_provider")
def test_max_tokens_propagation_relations(self, mock_create_provider):
"""Test that max_tokens is passed to generate_typed in extract_relations_llm."""
# Setup mock
mock_llm = MagicMock()
mock_create_provider.return_value = mock_llm
mock_llm.is_available.return_value = True
# Setup return value to avoid pydantic validation errors
mock_response = MagicMock()
mock_response.relations = []
mock_llm.generate_typed.return_value = mock_response
# Create dummy entities
entities = [Entity(text="Foo", label="ORG", start_char=0, end_char=3)]
# Call the function with max_tokens
extract_relations_llm(
text="some text",
entities=entities,
provider="openai",
model="gpt-4",
max_tokens=128000
)
# Check if generate_typed was called with max_tokens
args, kwargs = mock_llm.generate_typed.call_args
print(f"Relations Call kwargs: {kwargs}")
self.assertIn("max_tokens", kwargs)
self.assertEqual(kwargs["max_tokens"], 128000)
@patch("semantica.semantic_extract.methods.create_provider")
def test_max_tokens_propagation_entities(self, mock_create_provider):
"""Test that max_tokens is passed to generate_typed in extract_entities_llm."""
# Setup mock
mock_llm = MagicMock()
mock_create_provider.return_value = mock_llm
mock_llm.is_available.return_value = True
# Setup return value to avoid pydantic validation errors
mock_response = MagicMock()
mock_response.entities = []
mock_llm.generate_typed.return_value = mock_response
# Call the function with max_tokens
extract_entities_llm(
text="some text",
provider="openai",
model="gpt-4",
max_tokens=128000
)
# Check if generate_typed was called with max_tokens
args, kwargs = mock_llm.generate_typed.call_args
print(f"Entities Call kwargs: {kwargs}")
self.assertIn("max_tokens", kwargs)
self.assertEqual(kwargs["max_tokens"], 128000)
@patch("semantica.semantic_extract.methods.create_provider")
def test_max_tokens_propagation_triplets(self, mock_create_provider):
"""Test that max_tokens is passed to generate_typed in extract_triplets_llm."""
# Setup mock
mock_llm = MagicMock()
mock_create_provider.return_value = mock_llm
mock_llm.is_available.return_value = True
# Setup return value to avoid pydantic validation errors
mock_response = MagicMock()
mock_response.triplets = []
mock_llm.generate_typed.return_value = mock_response
# Call the function with max_tokens
extract_triplets_llm(
text="some text",
provider="openai",
model="gpt-4",
max_tokens=128000
)
# Check if generate_typed was called with max_tokens
args, kwargs = mock_llm.generate_typed.call_args
print(f"Triplets Call kwargs: {kwargs}")
self.assertIn("max_tokens", kwargs)
self.assertEqual(kwargs["max_tokens"], 128000)
if __name__ == "__main__":
unittest.main()
View File
@@ -0,0 +1,125 @@
import statistics
import time
import os
import sys
from typing import Dict, List
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../..")))
from semantica.semantic_extract.event_detector import EventDetector
from semantica.semantic_extract.ner_extractor import NERExtractor
from semantica.semantic_extract.relation_extractor import RelationExtractor
from semantica.semantic_extract.semantic_analyzer import SemanticAnalyzer
from semantica.semantic_extract.semantic_network_extractor import SemanticNetworkExtractor
from semantica.semantic_extract.triplet_extractor import TripletExtractor
from semantica.utils.progress_tracker import get_progress_tracker
def _make_documents(n: int) -> List[Dict[str, str]]:
base = (
"Apple Inc. was founded by Steve Jobs in 1976 and is headquartered in Cupertino, California. "
"Microsoft Corporation was founded by Bill Gates and Paul Allen in 1975. "
"In 2014, Apple acquired Beats Electronics for $3 billion. "
"In 2023, Google announced a partnership with OpenAI to improve search experiences."
)
return [{"id": f"doc_{i}", "content": f"{base} Document number {i}."} for i in range(n)]
def _median_seconds(fn, repeats: int = 3) -> float:
times = []
for _ in range(repeats):
start = time.perf_counter()
fn()
times.append(time.perf_counter() - start)
return statistics.median(times)
def _bench(label: str, fn, repeats: int = 3) -> dict:
fn()
seconds = _median_seconds(fn, repeats=repeats)
return {"label": label, "seconds": seconds}
def main():
progress = get_progress_tracker()
progress.displays = []
docs = _make_documents(80)
texts = [d["content"] for d in docs]
ner = NERExtractor(method="pattern")
rel = RelationExtractor(method="pattern")
trip = TripletExtractor(method="pattern")
events = EventDetector(method="pattern")
analyzer = SemanticAnalyzer()
net = SemanticNetworkExtractor(ner_method="pattern", relation_method="pattern")
results = []
def ner_parallel():
ner.extract(texts)
def ner_seq():
ner.extract(texts, max_workers=1)
results.append(_bench("NER batch (default workers)", ner_parallel))
results.append(_bench("NER batch (max_workers=1)", ner_seq))
entities_batch = ner.extract(texts)
def rel_parallel():
rel.extract(texts, entities_batch)
def rel_seq():
rel.extract(texts, entities_batch, max_workers=1)
results.append(_bench("Relation batch (default workers)", rel_parallel))
results.append(_bench("Relation batch (max_workers=1)", rel_seq))
def trip_parallel():
trip.extract(texts)
def trip_seq():
trip.extract(texts, max_workers=1)
results.append(_bench("Triplet pipeline (default workers)", trip_parallel))
results.append(_bench("Triplet pipeline (max_workers=1)", trip_seq))
def ev_parallel():
events.detect_events(texts)
def ev_seq():
events.detect_events(texts, max_workers=1)
results.append(_bench("Event detection (default workers)", ev_parallel))
results.append(_bench("Event detection (max_workers=1)", ev_seq))
def analyzer_parallel():
analyzer.analyze(texts)
def analyzer_seq():
analyzer.analyze(texts, max_workers=1)
results.append(_bench("Semantic analysis (default workers)", analyzer_parallel))
results.append(_bench("Semantic analysis (max_workers=1)", analyzer_seq))
def net_parallel():
net.extract_network(texts)
def net_seq():
net.extract_network(texts, max_workers=1)
results.append(_bench("Semantic network (default workers)", net_parallel))
results.append(_bench("Semantic network (max_workers=1)", net_seq))
per_doc = []
for row in results:
per_doc.append({**row, "ms_per_doc": (row["seconds"] / len(texts)) * 1000.0})
print(f"Documents: {len(texts)}")
for row in per_doc:
print(f"{row['label']}: {row['seconds']:.3f}s ({row['ms_per_doc']:.2f} ms/doc)")
if __name__ == "__main__":
main()
+84
View File
@@ -7,8 +7,12 @@ import os
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../..")))
from semantica.semantic_extract.ner_extractor import NERExtractor
from semantica.semantic_extract.ner_extractor import Entity as NEREntity
from semantica.semantic_extract.relation_extractor import RelationExtractor
from semantica.semantic_extract.triplet_extractor import TripletExtractor
from semantica.semantic_extract.event_detector import EventDetector
from semantica.semantic_extract.semantic_analyzer import SemanticAnalyzer
from semantica.semantic_extract.semantic_network_extractor import SemanticNetworkExtractor
from semantica.semantic_extract.named_entity_recognizer import Entity
from semantica.semantic_extract.relation_extractor import Relation
@@ -51,6 +55,21 @@ class TestExtractors(unittest.TestCase):
self.assertIsInstance(entities, list)
mock_get_method.assert_called()
@patch("semantica.semantic_extract.methods.get_entity_method")
def test_ner_extraction_batch_via_extract_entities(self, mock_get_method):
mock_method = MagicMock()
mock_method.extract_entities.return_value = []
mock_get_method.return_value = mock_method
extractor = NERExtractor(method="pattern")
results = extractor.extract_entities(
["Test text 1", "Test text 2"],
)
self.assertIsInstance(results, list)
self.assertEqual(len(results), 2)
self.assertTrue(all(isinstance(r, list) for r in results))
@patch("semantica.semantic_extract.methods.get_relation_method")
def test_relation_extraction(self, mock_get_method):
"""Test relation extraction call"""
@@ -65,6 +84,21 @@ class TestExtractors(unittest.TestCase):
self.assertIsInstance(relations, list)
mock_get_method.assert_called()
@patch("semantica.semantic_extract.methods.get_relation_method")
def test_relation_extraction_batch_via_extract_relations(self, mock_get_method):
mock_method = MagicMock()
mock_method.extract_relations.return_value = []
mock_get_method.return_value = mock_method
extractor = RelationExtractor(method="pattern")
texts = ["A knows B", "A knows B"]
entities = [NEREntity(text="A", label="PERSON", start_char=0, end_char=1, confidence=1.0)]
results = extractor.extract_relations(texts, entities)
self.assertIsInstance(results, list)
self.assertEqual(len(results), 2)
self.assertTrue(all(isinstance(r, list) for r in results))
@patch("semantica.semantic_extract.methods.get_triplet_method")
def test_triplet_extraction(self, mock_get_method):
"""Test triplet extraction call"""
@@ -81,5 +115,55 @@ class TestExtractors(unittest.TestCase):
self.assertIsInstance(triplets, list)
mock_get_method.assert_called()
@patch("semantica.semantic_extract.methods.get_triplet_method")
def test_triplet_extraction_batch_via_extract_triplets(self, mock_get_method):
mock_method = MagicMock()
mock_method.extract_triplets.return_value = []
mock_get_method.return_value = mock_method
extractor = TripletExtractor(method="pattern")
texts = ["A knows A", "A knows A"]
entities_batch = [[NEREntity(text="A", label="PERSON", start_char=0, end_char=1, confidence=1.0)] for _ in texts]
relations_batch = [[] for _ in texts]
results = extractor.extract_triplets(
texts,
entities=entities_batch,
relations=relations_batch,
)
self.assertIsInstance(results, list)
self.assertEqual(len(results), 2)
self.assertTrue(all(isinstance(r, list) for r in results))
def test_event_detector_batch_via_detect_events(self):
detector = EventDetector()
texts = ["Apple acquired Beats in 2014.", "Google announced a partnership in 2023."]
results = detector.detect_events(texts)
self.assertIsInstance(results, list)
self.assertEqual(len(results), 2)
self.assertTrue(all(isinstance(r, list) for r in results))
def test_semantic_analyzer_batch_parallel(self):
analyzer = SemanticAnalyzer()
texts = ["A short sentence.", "Another short sentence."]
results = analyzer.analyze(texts)
self.assertIsInstance(results, list)
self.assertEqual(len(results), 2)
self.assertTrue(all(isinstance(r, dict) for r in results))
def test_semantic_network_batch_via_extract_network(self):
extractor = SemanticNetworkExtractor()
texts = ["A knows B.", "C knows D."]
entities_batch = [[] for _ in texts]
relations_batch = [[] for _ in texts]
results = extractor.extract_network(
texts,
entities=entities_batch,
relations=relations_batch,
)
self.assertIsInstance(results, list)
self.assertEqual(len(results), 2)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,378 @@
import unittest
import time
import os
from dotenv import load_dotenv
load_dotenv()
from semantica.semantic_extract.ner_extractor import NERExtractor
from semantica.semantic_extract.relation_extractor import RelationExtractor
from semantica.semantic_extract.triplet_extractor import TripletExtractor
from semantica.semantic_extract.event_detector import EventDetector
from semantica.semantic_extract.semantic_network_extractor import SemanticNetworkExtractor
from semantica.semantic_extract.methods import _result_cache
class TestGroqRealWorldPerformance(unittest.TestCase):
"""
Real-world performance test suite using Groq LLM.
Tests parallel processing, caching, and correctness.
"""
@classmethod
def setUpClass(cls):
cls.api_key = os.getenv("GROQ_API_KEY")
if not cls.api_key:
raise unittest.SkipTest("GROQ_API_KEY is not set")
cls.metrics_file = os.path.join(os.getcwd(), "groq_metrics.txt")
# Real-world sample texts (mix of tech, business, and general)
cls.sample_texts = [
"""Apple Inc. is planning to launch a new AI-powered iPhone in late 2024.
CEO Tim Cook announced that the device will feature a neural engine capable of
processing 50 trillion operations per second. The company's stock rose 5% following the news.""",
"""Microsoft Corporation has acquired Activision Blizzard for $68.7 billion.
Satya Nadella, Microsoft's Chairman and CEO, stated that this acquisition will
accelerate growth in Microsoft's gaming business across mobile, PC, console, and cloud.""",
"""Elon Musk's SpaceX successfully launched the Starship rocket from Boca Chica, Texas.
The mission aims to test new heat shield technology essential for future Mars missions.
NASA Administrator Bill Nelson congratulated the team on the achievement.""",
"""Google DeepMind introduced Gemini, a new multimodal AI model.
Sundar Pichai emphasized that Gemini represents a significant leap forward in
AI capabilities, outperforming GPT-4 on several benchmarks including MMLU.""",
"""Amazon Web Services (AWS) announced a partnership with Anthropic to develop
reliable and high-performance foundation models. Amazon is investing up to $4 billion
in the AI safety startup founded by Dario Amodei."""
]
# Warm up: Ensure modules are loaded
print("\n[Setup] Initializing extractors...")
cls.extractor = NERExtractor(
method="llm",
provider="groq",
llm_model="llama-3.3-70b-versatile",
api_key=cls.api_key,
)
cls.relation_extractor = RelationExtractor(
method="llm",
provider="groq",
llm_model="llama-3.3-70b-versatile",
api_key=cls.api_key,
)
cls.triplet_extractor = TripletExtractor(
method="llm",
provider="groq",
llm_model="llama-3.3-70b-versatile",
api_key=cls.api_key,
)
cls.event_detector = EventDetector(
method="llm",
provider="groq",
llm_model="llama-3.3-70b-versatile",
api_key=cls.api_key,
)
cls.network_extractor = SemanticNetworkExtractor(
method="llm",
provider="groq",
llm_model="llama-3.3-70b-versatile",
api_key=cls.api_key,
)
def setUp(self):
# Clear cache before specific performance tests to ensure fair comparison
# (Unless testing cache specifically)
if _result_cache:
_result_cache._caches["entities"].clear()
_result_cache._caches["relations"].clear()
def log_metrics(self, message):
print(message)
with open(self.metrics_file, "a") as f:
f.write(message + "\n")
f.flush()
def test_01_parallel_vs_sequential_performance(self):
"""Compare sequential vs parallel extraction speed."""
try:
self.log_metrics("\n" + "="*60)
self.log_metrics("TEST 1: Sequential vs Parallel Processing Performance")
self.log_metrics("="*60)
extractor = NERExtractor(method="llm", provider="groq", api_key=self.api_key, model="llama-3.3-70b-versatile")
# 1. Sequential Run (Max workers = 1)
self.log_metrics("\nStarting Sequential Extraction (5 documents)...")
start_time = time.time()
seq_results = extractor.extract(self.sample_texts, max_workers=1)
seq_time = time.time() - start_time
self.log_metrics(f"Sequential Time: {seq_time:.4f}s")
self.log_metrics(f"Average Latency: {seq_time/len(self.sample_texts):.4f}s per doc")
# Clear cache to force re-extraction for parallel test
_result_cache._caches["entities"].clear()
# 2. Parallel Run (Max workers = 5)
self.log_metrics("\nStarting Parallel Extraction (5 documents, 5 workers)...")
start_time = time.time()
par_results = extractor.extract(self.sample_texts, max_workers=5)
par_time = time.time() - start_time
self.log_metrics(f"Parallel Time: {par_time:.4f}s")
self.log_metrics(f"Average Latency: {par_time/len(self.sample_texts):.4f}s per doc")
# Analysis
speedup = seq_time / par_time if par_time > 0 else 0
self.log_metrics(f"\n>>> Performance Gain: {speedup:.2f}x speedup")
self.log_metrics(f">>> Latency Reduction: {(seq_time - par_time):.4f}s total time saved")
self.assertLess(par_time, seq_time * 1.35, "Parallel processing should not be significantly slower")
self.assertEqual(len(seq_results), len(self.sample_texts))
self.assertEqual(len(par_results), len(self.sample_texts))
except Exception as e:
self.log_metrics(f"ERROR in Test 1: {e}")
raise
def test_02_caching_latency_reduction(self):
"""Measure latency reduction from caching."""
try:
self.log_metrics("\n" + "="*60)
self.log_metrics("TEST 2: Caching Performance & Latency Reduction")
self.log_metrics("="*60)
extractor = NERExtractor(method="llm", provider="groq", api_key=self.api_key, model="llama-3.3-70b-versatile")
text = [self.sample_texts[0]]
# 1. Cold Cache
_result_cache._caches["entities"].clear()
self.log_metrics("\nCold Cache Request...")
start_time = time.time()
extractor.extract(text)
cold_time = time.time() - start_time
self.log_metrics(f"Cold Cache Time: {cold_time:.4f}s")
cache_size_after_cold = _result_cache.get_stats()["entities"]["size"]
# 2. Warm Cache
self.log_metrics("\nWarm Cache Request (Identical Query)...")
start_time = time.time()
extractor.extract(text)
warm_time = time.time() - start_time
self.log_metrics(f"Warm Cache Time: {warm_time:.6f}s")
cache_size_after_warm = _result_cache.get_stats()["entities"]["size"]
# Analysis
reduction = (cold_time - warm_time) / cold_time * 100
self.log_metrics(f"\n>>> Latency Reduction: {reduction:.2f}%")
self.assertLess(warm_time, 1.0, "Warm cache response should be fast (<1.0s)")
# self.assertGreater(reduction, 50, "Caching should reduce latency by >50%")
if reduction < 30:
self.log_metrics(f"WARNING: Caching reduction is low ({reduction:.2f}%)")
self.assertGreater(reduction, 20, "Caching should reduce latency by >20%")
self.assertGreater(cache_size_after_cold, 0, "Cache should store entity results")
self.assertEqual(cache_size_after_warm, cache_size_after_cold, "Warm request should hit the cache")
except Exception as e:
self.log_metrics(f"ERROR in Test 2: {e}")
raise
def test_03_correctness_and_entity_matching(self):
"""Verify extraction correctness and data quality."""
try:
self.log_metrics("\n" + "="*60)
self.log_metrics("TEST 3: Extraction Correctness & Data Quality")
self.log_metrics("="*60)
# Use a specific text with clear entities
text = "Satya Nadella is the CEO of Microsoft."
extractor = NERExtractor(method="llm", provider="groq", api_key=self.api_key, model="llama-3.3-70b-versatile")
entities = extractor.extract([text])[0] # List of lists
self.log_metrics(f"\nInput: {text}")
self.log_metrics(f"Extracted Entities: {[e.text + '(' + e.label + ')' for e in entities]}")
# Validation
found_person = any(e.label == "PERSON" and "Satya" in e.text for e in entities)
found_org = any(e.label == "ORG" and "Microsoft" in e.text for e in entities)
if not found_org:
self.log_metrics("FAILURE: Did not find Microsoft as ORG. Found entities:")
for e in entities:
self.log_metrics(f" - {e.text}: {e.label}")
self.assertTrue(found_person, "Failed to extract Satya Nadella as PERSON")
self.assertTrue(found_org, "Failed to extract Microsoft as ORG")
self.log_metrics("\n>>> Correctness Verification: PASS")
self.log_metrics(" - Identified PERSON entity")
self.log_metrics(" - Identified ORG entity")
self.log_metrics(" - Pydantic models validated successfully")
except Exception as e:
self.log_metrics(f"ERROR in Test 3: {e}")
raise
def test_4_relation_extraction(self):
"""Test Relation Extraction capabilities"""
print("\n" + "="*60)
print("TEST 4: Relation Extraction")
print("="*60)
text = self.sample_texts[1] # Microsoft acquisition text
print(f"\nInput: {text[:100]}...")
# First extract entities
entities = self.__class__.extractor.extract_entities(text)
self.assertTrue(len(entities) > 0, "Should extract entities first")
# Extract relations
start_time = time.time()
relations = self.__class__.relation_extractor.extract_relations(text, entities)
duration = time.time() - start_time
print(f"Extracted {len(relations)} relations in {duration:.4f}s")
for r in relations:
print(f" - {r.subject.text} -> {r.predicate} -> {r.object.text}")
self.assertTrue(len(relations) > 0, "Should extract relations")
# Verify specific relation (Microsoft -> acquired -> Activision Blizzard)
found_acquisition = False
for r in relations:
if "Microsoft" in r.subject.text and "Activision" in r.object.text:
found_acquisition = True
break
if not found_acquisition:
# Fallback check - sometimes subject/object might be swapped or different wording
for r in relations:
if "Activision" in r.subject.text and "Microsoft" in r.object.text:
found_acquisition = True
break
self.assertTrue(found_acquisition, "Should find acquisition relation between Microsoft and Activision")
def test_5_triplet_extraction(self):
"""Test RDF Triplet Extraction capabilities"""
print("\n" + "="*60)
print("TEST 5: Triplet Extraction")
print("="*60)
text = self.sample_texts[0] # Apple text
print(f"\nInput: {text[:100]}...")
# Pipeline: Entities -> Relations -> Triplets
entities = self.__class__.extractor.extract_entities(text)
relations = self.__class__.relation_extractor.extract_relations(text, entities)
start_time = time.time()
triplets = self.__class__.triplet_extractor.extract_triplets(text, entities, relations)
duration = time.time() - start_time
print(f"Extracted {len(triplets)} triplets in {duration:.4f}s")
for t in triplets:
print(f" - <{t.subject}> <{t.predicate}> <{t.object}>")
self.assertTrue(len(triplets) > 0, "Should extract triplets")
# Check for Apple related triplet
found_apple = False
for t in triplets:
if "Apple" in t.subject or "Apple" in t.object:
found_apple = True
break
self.assertTrue(found_apple, "Should find Apple-related triplet")
def test_6_event_detection(self):
"""Test Event Detection capabilities"""
print("\n" + "="*60)
print("TEST 6: Event Detection")
print("="*60)
text = self.sample_texts[2] # SpaceX launch text
print(f"\nInput: {text[:100]}...")
start_time = time.time()
events = self.__class__.event_detector.detect_events(text)
duration = time.time() - start_time
print(f"Detected {len(events)} events in {duration:.4f}s")
for e in events:
print(f" - [{e.event_type}] {e.text} (Participants: {e.participants})")
self.assertTrue(len(events) > 0, "Should detect events")
# Verify launch event
found_launch = False
for e in events:
if "launch" in e.event_type.lower() or "launch" in e.text.lower():
found_launch = True
break
self.assertTrue(found_launch, "Should detect launch event")
def test_7_semantic_network(self):
"""Test Semantic Network Extraction capabilities"""
print("\n" + "="*60)
print("TEST 7: Semantic Network Extraction")
print("="*60)
text = self.sample_texts[3] # Google DeepMind text
print(f"\nInput: {text[:100]}...")
# Extract base components first
entities = self.__class__.extractor.extract_entities(text)
relations = self.__class__.relation_extractor.extract_relations(text, entities)
start_time = time.time()
network = self.__class__.network_extractor.extract_network(text, entities=entities, relations=relations)
duration = time.time() - start_time
print(f"Extracted Network in {duration:.4f}s")
print(f" - Nodes: {len(network.nodes)}")
print(f" - Edges: {len(network.edges)}")
self.assertTrue(len(network.nodes) > 0, "Should have nodes")
self.assertTrue(len(network.edges) > 0, "Should have edges")
# Verify Google/DeepMind/Gemini nodes exist
node_labels = [n.label for n in network.nodes]
print(f" - Node Labels: {node_labels}")
self.assertTrue(any("Gemini" in l for l in node_labels), "Should contain Gemini node")
def test_8_parallel_event_detection(self):
"""Test Parallel Event Detection capabilities"""
print("\n" + "="*60)
print("TEST 8: Parallel Event Detection")
print("="*60)
# Create a larger batch by duplicating sample texts
batch_texts = self.sample_texts * 2 # 10 documents
# 1. Sequential Run
print("\nStarting Sequential Event Detection (10 documents)...")
start_time = time.time()
seq_results = self.__class__.event_detector.extract(batch_texts, max_workers=1)
seq_time = time.time() - start_time
print(f"Sequential Time: {seq_time:.4f}s")
# 2. Parallel Run
print("\nStarting Parallel Event Detection (10 documents, 5 workers)...")
start_time = time.time()
par_results = self.__class__.event_detector.extract(batch_texts, max_workers=5)
par_time = time.time() - start_time
print(f"Parallel Time: {par_time:.4f}s")
# Analysis
speedup = seq_time / par_time if par_time > 0 else 0
print(f"\n>>> Performance Gain: {speedup:.2f}x speedup")
self.assertEqual(len(seq_results), len(batch_texts))
self.assertEqual(len(par_results), len(batch_texts))
# Verify results match (order should be preserved)
for i in range(len(batch_texts)):
self.assertEqual(len(seq_results[i]), len(par_results[i]), f"Result count mismatch at index {i}")
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,63 @@
import os
import pytest
try:
from dotenv import load_dotenv
load_dotenv()
except Exception:
pass
pytest.importorskip("groq")
from semantica.semantic_extract.ner_extractor import NERExtractor
from semantica.semantic_extract.relation_extractor import RelationExtractor
from semantica.semantic_extract.triplet_extractor import TripletExtractor
def test_groq_llm_smoke_entities_relations_triplets():
if not os.getenv("GROQ_API_KEY"):
pytest.skip("GROQ_API_KEY is not set")
text = (
"Apple acquired Beats in 2014 for $3 billion. "
"Steve Jobs founded Apple. "
"Beats is based in California."
)
model = "llama-3.3-70b-versatile"
entities = NERExtractor(method="llm").extract(
text,
provider="groq",
model=model,
temperature=0.0,
max_tokens=250,
)
assert isinstance(entities, list)
assert len(entities) > 0
assert len(entities) <= 30
relations = RelationExtractor(method="llm").extract(
text,
entities=entities,
provider="groq",
model=model,
temperature=0.0,
max_tokens=350,
max_entities_prompt=12,
)
assert isinstance(relations, list)
assert len(relations) <= 30
triplets = TripletExtractor(method="llm").extract(
text,
entities=entities,
relations=relations,
provider="groq",
model=model,
temperature=0.0,
max_tokens=350,
)
assert isinstance(triplets, list)
assert len(triplets) <= 40
+286
View File
@@ -0,0 +1,286 @@
import time
import unittest
print("Starting tests module...")
from unittest.mock import MagicMock, patch
from semantica.semantic_extract.providers import create_provider, ProviderPool, _provider_pool
from semantica.semantic_extract.ner_extractor import NERExtractor
from semantica.semantic_extract.relation_extractor import RelationExtractor
from semantica.semantic_extract.triplet_extractor import TripletExtractor, Triplet
from semantica.semantic_extract.methods import _result_cache, extract_entities_llm, extract_relations_llm, extract_triplets_llm, match_entity
from semantica.semantic_extract.ner_extractor import Entity
class TestSemanticExtractImprovements(unittest.TestCase):
def setUp(self):
_provider_pool.clear()
# Clear cache before each test
if _result_cache:
_result_cache._caches["entities"].clear()
_result_cache._caches["relations"].clear()
_result_cache._caches["triplets"].clear()
def test_entity_matching(self):
print("\nTesting Entity Matching...")
entities = [
Entity(text="Apple Inc.", label="ORG", start_char=0, end_char=10, confidence=1.0),
Entity(text="Steve Jobs", label="PERSON", start_char=0, end_char=10, confidence=1.0)
]
# Exact match
m1 = match_entity("Apple Inc.", entities)
self.assertIsNotNone(m1)
self.assertEqual(m1.text, "Apple Inc.")
# Case insensitive
m2 = match_entity("apple inc.", entities)
self.assertIsNotNone(m2)
self.assertEqual(m2.text, "Apple Inc.")
# Substring/Partial match (should work via calculate_similarity)
# "Apple" is contained in "Apple Inc."
# calculate_similarity gives a boost for containment
m3 = match_entity("Apple", entities)
if m3:
self.assertEqual(m3.text, "Apple Inc.")
print(" Partial match 'Apple' -> 'Apple Inc.' successful.")
else:
print(" Partial match 'Apple' -> 'Apple Inc.' failed (score too low).")
# No match
m4 = match_entity("Microsoft", entities)
self.assertIsNone(m4)
print(" No match verified.")
# Synonym match
# We need entities that match the synonym keys in methods.py (e.g. "acquired" -> "bought")
# Let's create an entity "bought"
rel_entities = [Entity(text="bought", label="RELATION", start_char=0, end_char=6, confidence=1.0)]
m5 = match_entity("acquired", rel_entities)
self.assertIsNotNone(m5)
self.assertEqual(m5.text, "bought")
print(" Synonym match 'acquired' -> 'bought' verified.")
# Empty input
m6 = match_entity("", entities)
self.assertIsNone(m6)
print(" Empty input handled.")
def test_caching(self):
print("\nTesting Caching...")
text = "Apple Inc. was founded in 1976."
# Mock provider
mock_provider = MagicMock()
mock_provider.is_available.return_value = True
# Setup mock response for entities
mock_entities_response = MagicMock()
mock_entities_response.entities = [
MagicMock(text="Apple Inc.", label="ORG", confidence=0.9),
MagicMock(text="1976", label="DATE", confidence=0.9)
]
mock_provider.generate_typed.return_value = mock_entities_response
with patch('semantica.semantic_extract.methods.create_provider', return_value=mock_provider) as mock_create:
# First call - should hit provider
print(" First call (cache miss)...")
results1 = extract_entities_llm(text, provider="openai", model="gpt-4", api_key="test")
self.assertEqual(len(results1), 2)
self.assertEqual(mock_provider.generate_typed.call_count, 1)
# Check cache state
print(f" Cache size: {len(_result_cache._caches['entities'])}")
# Second call - should hit cache
print(" Second call (cache hit)...")
results2 = extract_entities_llm(text, provider="openai", model="gpt-4", api_key="test")
self.assertEqual(len(results2), 2)
# Provider should NOT be called again
self.assertEqual(mock_provider.generate_typed.call_count, 1)
print(" Cache hit verified for entities.")
def test_secure_caching(self):
"""Test that sensitive parameters are excluded from cache keys."""
print("\nTesting Secure Caching...")
text = "Security test."
# Mock provider
mock_provider = MagicMock()
mock_provider.is_available.return_value = True
mock_entities_response = MagicMock()
mock_entities_response.entities = [MagicMock(text="Test", label="TEST", confidence=1.0)]
mock_provider.generate_typed.return_value = mock_entities_response
with patch('semantica.semantic_extract.methods.create_provider', return_value=mock_provider):
# First call with one API key
extract_entities_llm(text, provider="openai", model="gpt-4", api_key="secret_key_1")
# Second call with DIFFERENT API key
# If secure caching is working, this should be a CACHE HIT because api_key is ignored
extract_entities_llm(text, provider="openai", model="gpt-4", api_key="secret_key_2")
# Provider should have been called ONLY ONCE
self.assertEqual(mock_provider.generate_typed.call_count, 1)
print(" Secure caching verified: Changing API key did not trigger new extraction.")
# Verify cache content
self.assertIn("entities", _result_cache._caches)
self.assertTrue(len(_result_cache._caches["entities"]) > 0)
def test_provider_pool(self):
print("\nTesting Provider Pool...")
# Create provider twice with same args
# We need to mock the actual provider init to avoid API keys requirement if not present
with patch('semantica.semantic_extract.providers.OpenAIProvider') as MockProvider:
MockProvider.side_effect = lambda *args, **kwargs: MagicMock()
p1 = create_provider("openai", api_key="test", model_name="gpt-4")
p2 = create_provider("openai", api_key="test", model_name="gpt-4")
# Should be same instance
self.assertIs(p1, p2)
print(" Provider reuse verified.")
# Different args
p3 = create_provider("openai", api_key="test", model_name="gpt-3.5")
self.assertIsNot(p1, p3)
print(" Different args create new instance verified.")
# Explicitly not using pool
p4 = create_provider("openai", use_pool=False, api_key="test", model_name="gpt-4")
self.assertIsNot(p1, p4)
print(" Opt-out of pool verified.")
def test_ner_parallel_processing(self):
print("\nTesting NER Parallel Processing...")
extractor = NERExtractor(method="pattern") # Use pattern which is fast/local
# Mock extract_entities to simulate work and track thread execution
original_extract = extractor.extract_entities
def mock_extract(text, **kwargs):
time.sleep(0.1) # Simulate delay
return original_extract(text, **kwargs)
extractor.extract_entities = mock_extract
texts = ["Text 1", "Text 2", "Text 3", "Text 4"]
start_time = time.time()
results = extractor.extract(texts)
end_time = time.time()
duration = end_time - start_time
print(f" Parallel NER (default workers) took {duration:.4f}s")
self.assertEqual(len(results), 4)
# Verify sequential fallback
start_time_seq = time.time()
extractor.extract(texts, max_workers=1)
end_time_seq = time.time()
duration_seq = end_time_seq - start_time_seq
print(f" Sequential NER took {duration_seq:.4f}s")
# Check if parallel was indeed parallel (faster)
# With 0.1s sleep * 4 items:
# Sequential ~ 0.4s
# Parallel (2 workers) ~ 0.2s + overhead
self.assertLess(duration, duration_seq * 0.8)
print(" Parallel execution speedup verified.")
def test_relation_parallel_processing(self):
print("\nTesting Relation Parallel Processing...")
extractor = RelationExtractor(method="pattern")
# Mock extract_relations
original_extract = extractor.extract_relations
def mock_extract(text, entities, **kwargs):
time.sleep(0.1)
return original_extract(text, entities, **kwargs)
extractor.extract_relations = mock_extract
texts = ["Text 1", "Text 2", "Text 3", "Text 4"]
entities = [[], [], [], []]
start_time = time.time()
results = extractor.extract(texts, entities)
end_time = time.time()
duration = end_time - start_time
print(f" Parallel RE (default workers) took {duration:.4f}s")
self.assertEqual(len(results), 4)
# Sequential
start_time_seq = time.time()
extractor.extract(texts, entities, max_workers=1)
end_time_seq = time.time()
duration_seq = end_time_seq - start_time_seq
print(f" Sequential RE took {duration_seq:.4f}s")
self.assertLess(duration, duration_seq * 0.8)
print(" Parallel execution speedup verified.")
def test_relation_extraction_fuzzy_matching(self):
print("\nTesting Relation Extraction Fuzzy Matching...")
extractor = RelationExtractor(method="pattern")
# Entities have formal names
entities = [
Entity(text="Apple Inc.", label="ORG", start_char=0, end_char=10, confidence=1.0),
Entity(text="Steve Jobs", label="PERSON", start_char=21, end_char=31, confidence=1.0)
]
# Text uses informal name "Apple"
text = "Apple was founded by Steve Jobs."
relations = extractor.extract(text, entities)
found = False
for rel in relations:
# Check if subject matches "Apple Inc." even though text said "Apple"
if rel.subject.text == "Apple Inc." and rel.object.text == "Steve Jobs" and rel.predicate == "founded_by":
found = True
print(" Successfully matched 'Apple' -> 'Apple Inc.' in relation extraction.")
break
self.assertTrue(found, "Failed to extract relation with fuzzy entity matching")
def test_triplet_parallel_processing(self):
print("\nTesting Triplet Parallel Processing...")
extractor = TripletExtractor(method="pattern")
# Mock extract_triplets
original_extract = extractor.extract_triplets
def mock_extract(text, **kwargs):
time.sleep(0.1)
# Return dummy triplets to avoid actual extraction overhead
return [Triplet(subject="s", predicate="p", object="o")]
extractor.extract_triplets = mock_extract
texts = ["Text 1", "Text 2", "Text 3", "Text 4"]
start_time = time.time()
results = extractor.extract(texts)
end_time = time.time()
duration = end_time - start_time
print(f" Parallel TE (default workers) took {duration:.4f}s")
self.assertEqual(len(results), 4)
# Sequential
start_time_seq = time.time()
extractor.extract(texts, max_workers=1)
end_time_seq = time.time()
duration_seq = end_time_seq - start_time_seq
print(f" Sequential TE took {duration_seq:.4f}s")
self.assertLess(duration, duration_seq * 0.8)
print(" Parallel execution speedup verified.")
if __name__ == "__main__":
suite = unittest.TestLoader().loadTestsFromTestCase(TestSemanticExtractImprovements)
unittest.TextTestRunner(verbosity=2).run(suite)
@@ -0,0 +1,157 @@
import unittest
from unittest.mock import patch, MagicMock
import sys
import os
# Ensure we test the local code, not the installed package
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '../../')))
from semantica.semantic_extract.ner_extractor import Entity
from semantica.semantic_extract.relation_extractor import Relation
from semantica.semantic_extract.triplet_extractor import Triplet
from semantica.semantic_extract.methods import extract_entities_llm
# We import providers later inside tests to allow patching
class TestSemanticClasses:
"""Test that core semantic classes do not have hardcoded max lengths."""
def test_entity_no_max_length(self):
long_text = "a" * 10000
entity = Entity(text=long_text, label="TEST", start_char=0, end_char=10000)
assert entity.text == long_text
assert len(entity.text) == 10000
def test_relation_no_max_length(self):
long_text = "a" * 10000
e1 = Entity(text="s", label="S", start_char=0, end_char=1)
e2 = Entity(text="o", label="O", start_char=0, end_char=1)
relation = Relation(subject=e1, predicate=long_text, object=e2)
assert relation.predicate == long_text
def test_triplet_no_max_length(self):
long_text = "a" * 10000
triplet = Triplet(subject=long_text, predicate="r", object="t")
assert triplet.subject == long_text
class TestProviderLimits:
"""Test that providers pass through correct length parameters."""
def test_openai_max_completion_tokens(self):
from semantica.semantic_extract.providers import OpenAIProvider
# Patch _init_client to avoid real client creation and import issues
with patch.object(OpenAIProvider, '_init_client', return_value=None):
provider = OpenAIProvider(api_key="fake")
# Manually mock client
mock_client = MagicMock()
mock_response = MagicMock()
mock_response.choices[0].message.content = "result"
mock_client.chat.completions.create.return_value = mock_response
provider.client = mock_client
provider.generate("prompt", max_completion_tokens=12345, top_p=0.9)
call_kwargs = mock_client.chat.completions.create.call_args[1]
assert call_kwargs["max_completion_tokens"] == 12345
assert call_kwargs["top_p"] == 0.9
assert "max_tokens" not in call_kwargs
def test_anthropic_max_tokens_defaults(self):
from semantica.semantic_extract.providers import AnthropicProvider
with patch.object(AnthropicProvider, '_init_client', return_value=None):
provider = AnthropicProvider(api_key="fake")
mock_client = MagicMock()
mock_response = MagicMock()
mock_response.content = [MagicMock(text="result")]
mock_client.messages.create.return_value = mock_response
provider.client = mock_client
provider.generate("prompt")
# Verify default is 8192 (new limit)
call_kwargs = mock_client.messages.create.call_args[1]
assert call_kwargs["max_tokens"] == 8192
# Test override
provider.generate("prompt", max_tokens=9999)
call_kwargs = mock_client.messages.create.call_args[1]
assert call_kwargs["max_tokens"] == 9999
def test_groq_max_completion_tokens(self):
from semantica.semantic_extract.providers import GroqProvider
with patch.object(GroqProvider, '_init_client', return_value=None):
provider = GroqProvider(api_key="fake")
mock_client = MagicMock()
mock_response = MagicMock()
mock_response.choices[0].message.content = "result"
mock_client.chat.completions.create.return_value = mock_response
provider.client = mock_client
provider.generate("prompt", max_completion_tokens=5000)
# Verify
call_kwargs = mock_client.chat.completions.create.call_args[1]
assert call_kwargs["max_completion_tokens"] == 5000
def test_gemini_params(self):
from semantica.semantic_extract.providers import GeminiProvider
with patch.object(GeminiProvider, '_init_client', return_value=None):
provider = GeminiProvider(api_key="fake")
mock_model = MagicMock()
mock_response = MagicMock()
mock_response.text = "result"
mock_model.generate_content.return_value = mock_response
provider.client = mock_model
provider.generate("prompt", top_k=10, candidate_count=2)
# Verify
call_kwargs = mock_model.generate_content.call_args[1]
gen_config = call_kwargs["generation_config"]
assert gen_config["top_k"] == 10
assert gen_config["candidate_count"] == 2
class TestChunkingDefaults:
"""Test that chunking defaults have been increased."""
@patch("semantica.semantic_extract.methods.create_provider")
@patch("semantica.semantic_extract.methods._extract_entities_chunked")
def test_openai_chunking_limit(self, mock_chunked, mock_create_provider):
# Setup
mock_llm = MagicMock()
mock_llm.is_available.return_value = True
mock_create_provider.return_value = mock_llm
# Text length = 10000 (Greater than old 4000, less than new 64000)
long_text = "a" * 10000
# Call without explicit max_text_length
extract_entities_llm(long_text, provider="openai", api_key="fake")
# Should NOT call chunked extraction because default is now 64000
mock_chunked.assert_not_called()
@patch("semantica.semantic_extract.methods.create_provider")
@patch("semantica.semantic_extract.methods._extract_entities_chunked")
def test_groq_chunking_limit(self, mock_chunked, mock_create_provider):
# Setup
mock_llm = MagicMock()
mock_llm.is_available.return_value = True
mock_create_provider.return_value = mock_llm
# Text length = 10000 (Greater than old 8000, less than new 64000)
long_text = "a" * 10000
extract_entities_llm(long_text, provider="groq", api_key="fake")
# Should NOT call chunked extraction because default is now 64000
mock_chunked.assert_not_called()
@@ -0,0 +1,123 @@
import multiprocessing
import time
from unittest.mock import MagicMock, patch
import pytest
from semantica.semantic_extract.config import resolve_max_workers
from semantica.semantic_extract.ner_extractor import Entity, NERExtractor
from semantica.semantic_extract.relation_extractor import RelationExtractor
from semantica.semantic_extract.triplet_extractor import TripletExtractor
from semantica.semantic_extract.semantic_network_extractor import SemanticNetworkExtractor
from semantica.semantic_extract.methods import filter_entities_for_text
from semantica.semantic_extract.schemas import RelationsResponse, RelationOut
def test_resolve_max_workers_defaults_and_clamps():
cpu_count = multiprocessing.cpu_count() or 1
assert resolve_max_workers(explicit=0) == 1
assert resolve_max_workers(explicit=-10) == 1
assert resolve_max_workers(explicit=1) == 1
assert resolve_max_workers(explicit=10**9) == min(cpu_count, 32)
assert resolve_max_workers(explicit=None, methods=["ml"]) == 1
def test_filter_entities_for_text_keeps_short_tokens():
text = "US AI lab in NY"
entities = [
Entity(text="US", label="GPE", start_char=0, end_char=2, confidence=1.0),
Entity(text="AI", label="TECH", start_char=3, end_char=5, confidence=1.0),
Entity(text="NY", label="GPE", start_char=13, end_char=15, confidence=1.0),
]
kept = filter_entities_for_text(text, entities, max_keep=2)
kept_texts = {e.text for e in kept}
assert "US" in kept_texts or "AI" in kept_texts or "NY" in kept_texts
def test_pattern_batch_defaults_to_single_worker_low_latency():
extractor = NERExtractor(method="pattern")
texts = [f"Text {i}" for i in range(8)]
extractor.extract(texts)
def test_relation_llm_prompt_filter_does_not_break_mapping():
entities = [Entity(text=f"VeryLongEntityName{i}", label="ORG", start_char=0, end_char=1, confidence=1.0) for i in range(120)]
ghost = Entity(text="Ghost", label="ORG", start_char=0, end_char=1, confidence=1.0)
entities.append(ghost)
captured = {}
class FakeLLM:
def is_available(self):
return True
def generate_typed(self, prompt, schema, **kwargs):
captured["prompt"] = prompt
return RelationsResponse(
relations=[
RelationOut(subject="Ghost", predicate="related_to", object="VeryLongEntityName0", confidence=0.9)
]
)
with patch("semantica.semantic_extract.methods.create_provider", return_value=FakeLLM()):
from semantica.semantic_extract.methods import extract_relations_llm
relations = extract_relations_llm(
"Short text mentioning VeryLongEntityName0 only.",
entities=entities,
provider="openai",
model="gpt-4",
max_entities_prompt=20,
)
assert "Ghost" not in captured["prompt"]
assert len(relations) == 1
assert relations[0].subject.text == "Ghost"
def test_triplet_extractor_reuses_sub_extractors():
ner_instance = MagicMock()
ner_instance.extract_entities.return_value = [
Entity(text="A", label="PERSON", start_char=0, end_char=1, confidence=1.0)
]
rel_instance = MagicMock()
rel_instance.extract_relations.return_value = []
ner_ctor = MagicMock(return_value=ner_instance)
rel_ctor = MagicMock(return_value=rel_instance)
with patch("semantica.semantic_extract.ner_extractor.NERExtractor", ner_ctor), patch(
"semantica.semantic_extract.relation_extractor.RelationExtractor", rel_ctor
), patch("semantica.semantic_extract.methods.get_triplet_method", return_value=lambda *args, **kwargs: []):
extractor = TripletExtractor(method="pattern")
extractor.extract_triplets("A text.")
extractor.extract_triplets("A text again.")
assert ner_ctor.call_count == 1
assert rel_ctor.call_count == 1
def test_semantic_network_extractor_reuses_sub_extractors():
ner_instance = MagicMock()
ner_instance.extract_entities.return_value = [
Entity(text="A", label="PERSON", start_char=0, end_char=1, confidence=1.0)
]
rel_instance = MagicMock()
rel_instance.extract_relations.return_value = []
ner_ctor = MagicMock(return_value=ner_instance)
rel_ctor = MagicMock(return_value=rel_instance)
with patch("semantica.semantic_extract.ner_extractor.NERExtractor", ner_ctor), patch(
"semantica.semantic_extract.relation_extractor.RelationExtractor", rel_ctor
):
extractor = SemanticNetworkExtractor(method="pattern")
extractor.extract_network("A text.")
extractor.extract_network("A text again.")
assert ner_ctor.call_count == 1
assert rel_ctor.call_count == 1
+118
View File
@@ -0,0 +1,118 @@
import unittest
from semantica.semantic_extract.methods import extract_relations_llm
from semantica.semantic_extract.ner_extractor import Entity
class FakeProvider:
def __init__(self, typed_payload=None, structured_payload=None):
self._typed_payload = typed_payload
self._structured_payload = structured_payload
def is_available(self):
return True
# Simulate typed output return: can be dict or an object with relations
def generate_typed(self, prompt, schema, **kwargs):
return self._typed_payload if self._typed_payload is not None else {"relations": []}
def generate_structured(self, prompt, **kwargs):
return self._structured_payload if self._structured_payload is not None else {"relations": []}
class TestLLMRelationExtraction(unittest.TestCase):
def setUp(self):
# Minimal realistic text and entities
self.text = "Apple reported revenue of $4.4 billion in Q1 2024."
self.entities = [
Entity(text="Apple", label="ORGANIZATION", start_char=0, end_char=5, confidence=0.99),
Entity(text="$4.4 billion", label="MONEY", start_char=26, end_char=39, confidence=0.99),
Entity(text="Q1 2024", label="DATE", start_char=43, end_char=51, confidence=0.99),
]
def _monkeypatch_provider(self, provider_instance):
# Monkeypatch create_provider used by extract_relations_llm
import semantica.semantic_extract.methods as methods
self._orig_create_provider = methods.create_provider
def _fake_create_provider(provider, model=None, **kwargs):
return provider_instance
methods.create_provider = _fake_create_provider
def tearDown(self):
# Restore original create_provider if patched
try:
import semantica.semantic_extract.methods as methods
if hasattr(self, "_orig_create_provider"):
methods.create_provider = self._orig_create_provider
except Exception:
pass
def test_typed_relations_parsed(self):
# Typed returns a dict compatible with parser
typed_payload = {
"relations": [
{
"subject": "Apple",
"predicate": "HAS_REVENUE",
"object": "$4.4 billion",
"confidence": 0.92,
},
{
"subject": "Apple",
"predicate": "IN_QUARTER",
"object": "Q1 2024",
"confidence": 0.9,
},
]
}
fake = FakeProvider(typed_payload=typed_payload)
self._monkeypatch_provider(fake)
rels = extract_relations_llm(
text=self.text,
entities=self.entities,
provider="groq",
model="llama-3.1-8b-instant",
relation_types=["HAS_REVENUE", "IN_QUARTER"],
verbose=True,
)
self.assertGreaterEqual(len(rels), 2, "Expected at least two relations from typed payload")
preds = {(r.subject.text, r.predicate, r.object.text) for r in rels}
self.assertIn(("Apple", "HAS_REVENUE", "$4.4 billion"), preds)
self.assertIn(("Apple", "IN_QUARTER", "Q1 2024"), preds)
def test_structured_fallback_used(self):
# Typed returns zero, structured has content
typed_payload = {"relations": []}
structured_payload = {
"relations": [
{
"subject": "Apple",
"predicate": "HAS_REVENUE",
"object": "$4.4 billion",
"confidence": 0.88,
}
]
}
fake = FakeProvider(typed_payload=typed_payload, structured_payload=structured_payload)
self._monkeypatch_provider(fake)
rels = extract_relations_llm(
text=self.text,
entities=self.entities,
provider="groq",
model="llama-3.1-8b-instant",
relation_types=["HAS_REVENUE"],
verbose=True,
)
self.assertEqual(len(rels), 1, "Expected fallback to structured JSON to yield one relation")
r = rels[0]
self.assertEqual(r.subject.text, "Apple")
self.assertEqual(r.object.text, "$4.4 billion")
self.assertEqual(r.predicate, "HAS_REVENUE")
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,149 @@
import unittest
import numpy as np
import time
import sys
import os
import logging
from unittest.mock import MagicMock, patch
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../..")))
from semantica.vector_store import VectorStore
from semantica.utils.exceptions import ProcessingError
class TestVectorStoreParallel(unittest.TestCase):
def setUp(self):
logging.getLogger("vector_store").setLevel(logging.ERROR)
self.dimension = 4
self.store = VectorStore(
backend="inmemory",
dimension=self.dimension,
)
self.store.embedder = MagicMock()
def test_embed_batch_success(self):
texts = ["a", "b", "c"]
expected_embeddings = [
np.array([0.1] * 4, dtype=np.float32),
np.array([0.2] * 4, dtype=np.float32),
np.array([0.3] * 4, dtype=np.float32),
]
self.store.embedder.generate_embeddings.return_value = expected_embeddings
results = self.store.embed_batch(texts)
self.assertEqual(len(results), 3)
self.assertTrue(np.allclose(results[0], expected_embeddings[0]))
self.store.embedder.generate_embeddings.assert_called_once_with(texts)
def test_embed_batch_fallback(self):
texts = ["a", "b"]
self.store.embedder.generate_embeddings.side_effect = Exception("Model error")
results = self.store.embed_batch(texts)
self.assertEqual(len(results), 2)
self.assertEqual(results[0].shape, (self.dimension,))
self.assertTrue(isinstance(results[0], np.ndarray))
def test_add_documents_empty(self):
ids = self.store.add_documents([])
self.assertEqual(ids, [])
def test_add_documents_metadata_mismatch(self):
with self.assertRaises(ValueError):
self.store.add_documents(["doc1"], metadata=[{}, {}])
def test_add_documents_parallel_success(self):
num_docs = 10
documents = [f"doc_{i}" for i in range(num_docs)]
metadata = [{"id": i} for i in range(num_docs)]
def mock_embed_batch(texts):
return [np.full(self.dimension, float(i)) for i, _ in enumerate(texts)]
with patch.object(self.store, "embed_batch", side_effect=mock_embed_batch):
ids = self.store.add_documents(
documents,
metadata,
batch_size=2,
parallel=True,
)
self.assertEqual(len(ids), num_docs)
self.assertEqual(len(self.store.vectors), num_docs)
for i, vec_id in enumerate(ids):
stored_meta = self.store.get_metadata(vec_id)
self.assertEqual(stored_meta["id"], i)
def test_add_documents_sequential_success(self):
num_docs = 5
documents = [f"doc_{i}" for i in range(num_docs)]
with patch.object(self.store, "embed_batch") as mock_batch:
mock_batch.return_value = [np.zeros(self.dimension) for _ in range(num_docs)]
ids = self.store.add_documents(documents, parallel=False)
self.assertEqual(len(ids), num_docs)
self.assertEqual(mock_batch.call_count, 1)
def test_add_documents_error_propagation(self):
documents = ["doc1", "doc2"]
with patch.object(self.store, "embed_batch", side_effect=ValueError("Embedding Error")):
with self.assertRaises(Exception):
self.store.add_documents(documents, parallel=True)
def test_performance_simulation(self):
num_batches = 4
batch_delay = 0.1
batch_size = 1
documents = [f"doc_{i}" for i in range(num_batches)]
def slow_embed(texts):
time.sleep(batch_delay)
return [np.zeros(self.dimension) for _ in texts]
with patch.object(self.store, "embed_batch", side_effect=slow_embed):
start_seq = time.time()
self.store.add_documents(documents, batch_size=batch_size, parallel=False)
dur_seq = time.time() - start_seq
self.store.vectors = {}
start_par = time.time()
self.store.add_documents(documents, batch_size=batch_size, parallel=True)
dur_par = time.time() - start_par
print(f"\nPerformance Test:")
print(f"Sequential Duration: {dur_seq:.4f}s")
print(f"Parallel Duration: {dur_par:.4f}s")
print(f"Speedup: {dur_seq / dur_par:.2f}x")
self.assertLess(dur_par, dur_seq * 0.7)
def test_add_documents_batch_size_edge_cases(self):
documents = ["a", "b", "c"]
with patch.object(self.store, "embed_batch") as mock_batch:
mock_batch.side_effect = lambda texts: [np.zeros(4) for _ in texts]
self.store.add_documents(documents, batch_size=100)
self.assertEqual(mock_batch.call_count, 1)
mock_batch.reset_mock()
self.store.add_documents(documents, batch_size=1)
self.assertEqual(mock_batch.call_count, 3)
if __name__ == "__main__":
unittest.main()