feat(classifier): add multilingual ECP inherence classifier POC
This commit is contained in:
+166
@@ -0,0 +1,166 @@
|
||||
"""Data models and validation schemas for Multilingual NLP Entity Inherence Classifier."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, field, asdict
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
|
||||
class DecisionCategory(str, Enum):
|
||||
DIRECT_INHERENT = "DIRECT_INHERENT"
|
||||
CONTEXTUAL_INHERENT = "CONTEXTUAL_INHERENT"
|
||||
TANGENTIAL = "TANGENTIAL"
|
||||
NOT_RELATED = "NOT_RELATED"
|
||||
|
||||
|
||||
class ErrorCode(str, Enum):
|
||||
INVALID_ECP_JSON = "invalid_ecp_json"
|
||||
INVALID_MARKDOWN = "invalid_markdown"
|
||||
UNSUPPORTED_LANGUAGE = "unsupported_language"
|
||||
EMPTY_CONTENT = "empty_content"
|
||||
MISSING_REQUIRED_FIELD = "missing_required_field"
|
||||
|
||||
|
||||
@dataclass
|
||||
class RelatedEntity:
|
||||
entity_id: str
|
||||
name: str
|
||||
relation_type: str
|
||||
weight: float
|
||||
aliases: List[str] = field(default_factory=list)
|
||||
scope: str = "general"
|
||||
confidence: float = 1.0
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: Dict[str, Any]) -> RelatedEntity:
|
||||
if not isinstance(data, dict):
|
||||
raise ValueError("Related entity must be a JSON object")
|
||||
|
||||
required = ["entity_id", "name", "relation_type", "weight"]
|
||||
for req in required:
|
||||
if req not in data or data[req] is None:
|
||||
raise ValueError(f"Missing required field in related entity: '{req}'")
|
||||
|
||||
return cls(
|
||||
entity_id=str(data["entity_id"]),
|
||||
name=str(data["name"]),
|
||||
relation_type=str(data["relation_type"]),
|
||||
weight=float(data["weight"]),
|
||||
aliases=list(data.get("aliases", [])),
|
||||
scope=str(data.get("scope", "general")),
|
||||
confidence=float(data.get("confidence", 1.0)),
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ECPSnapshot:
|
||||
target_entity_id: str
|
||||
target_name: str
|
||||
aliases: List[str]
|
||||
domain: str
|
||||
anchors: List[str]
|
||||
negative_anchors: List[str] = field(default_factory=list)
|
||||
graph_version: str = "1.0.0"
|
||||
related_entities: List[RelatedEntity] = field(default_factory=list)
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: Dict[str, Any]) -> ECPSnapshot:
|
||||
if not isinstance(data, dict):
|
||||
raise ValueError("ECP Snapshot payload must be a JSON object")
|
||||
|
||||
required_fields = ["target_entity_id", "target_name", "aliases", "domain", "anchors"]
|
||||
for field_name in required_fields:
|
||||
if field_name not in data or data[field_name] is None:
|
||||
raise ValueError(f"Missing required field in ECP Snapshot: '{field_name}'")
|
||||
|
||||
if not isinstance(data["aliases"], list):
|
||||
raise ValueError("Field 'aliases' must be a list of strings")
|
||||
if not isinstance(data["anchors"], list):
|
||||
raise ValueError("Field 'anchors' must be a list of strings")
|
||||
|
||||
neg_anchors = data.get("negative_anchors", [])
|
||||
if neg_anchors is not None and not isinstance(neg_anchors, list):
|
||||
raise ValueError("Field 'negative_anchors' must be a list of strings if provided")
|
||||
|
||||
related_data = data.get("related_entities", [])
|
||||
if related_data is not None and not isinstance(related_data, list):
|
||||
raise ValueError("Field 'related_entities' must be a list if provided")
|
||||
|
||||
related_objs = [RelatedEntity.from_dict(item) for item in (related_data or [])]
|
||||
|
||||
return cls(
|
||||
target_entity_id=str(data["target_entity_id"]),
|
||||
target_name=str(data["target_name"]),
|
||||
aliases=[str(a) for a in data["aliases"]],
|
||||
domain=str(data["domain"]),
|
||||
anchors=[str(a) for a in data["anchors"]],
|
||||
negative_anchors=[str(na) for na in (neg_anchors or [])],
|
||||
graph_version=str(data.get("graph_version", "1.0.0")),
|
||||
related_entities=related_objs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_json_str(cls, json_str: str) -> ECPSnapshot:
|
||||
try:
|
||||
data = json.loads(json_str)
|
||||
except Exception as e:
|
||||
raise ValueError(f"Invalid JSON format: {e}") from e
|
||||
return cls.from_dict(data)
|
||||
|
||||
|
||||
@dataclass
|
||||
class MatchedGraphEntity:
|
||||
entity_id: str
|
||||
name: str
|
||||
relation_type: str
|
||||
weight: float
|
||||
|
||||
|
||||
@dataclass
|
||||
class ClassificationResult:
|
||||
decision: DecisionCategory
|
||||
is_inherent: bool
|
||||
confidence: float
|
||||
detected_language: str
|
||||
matched_anchors: List[str] = field(default_factory=list)
|
||||
negative_matches: List[str] = field(default_factory=list)
|
||||
graph_matches: List[Dict[str, Any]] = field(default_factory=list)
|
||||
evidence: List[str] = field(default_factory=list)
|
||||
rationale: str = ""
|
||||
warnings: List[str] = field(default_factory=list)
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"decision": self.decision.value if isinstance(self.decision, DecisionCategory) else str(self.decision),
|
||||
"is_inherent": bool(self.is_inherent),
|
||||
"confidence": round(float(self.confidence), 4),
|
||||
"detected_language": str(self.detected_language),
|
||||
"matched_anchors": list(self.matched_anchors),
|
||||
"negative_matches": list(self.negative_matches),
|
||||
"graph_matches": list(self.graph_matches),
|
||||
"evidence": list(self.evidence),
|
||||
"rationale": str(self.rationale),
|
||||
"warnings": list(self.warnings),
|
||||
}
|
||||
|
||||
def to_json_str(self, indent: int = 2) -> str:
|
||||
return json.dumps(self.to_dict(), indent=indent, ensure_ascii=False)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ClassificationError:
|
||||
error_code: ErrorCode
|
||||
message: str
|
||||
details: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"error_code": self.error_code.value if isinstance(self.error_code, ErrorCode) else str(self.error_code),
|
||||
"message": str(self.message),
|
||||
"details": dict(self.details),
|
||||
}
|
||||
|
||||
def to_json_str(self, indent: int = 2) -> str:
|
||||
return json.dumps(self.to_dict(), indent=indent, ensure_ascii=False)
|
||||
Reference in New Issue
Block a user