| |
| """ |
| EvoRM Plugin: Evolvable Neuro-Symbolic Reasoning Framework |
| =========================================================== |
| A plug-and-play module that implements the EvoRM framework from TKDE paper. |
| Integrates with AdaCoAgentEA's LLM1_label_selector.py. |
| |
| Core components: |
| - RuleEncoding: distill LLM decisions into FOL rules |
| - HypergraphStorage: weighted hypergraph for rules + entity contexts |
| - TwoStageInferenceController: symbolic filtering + LLM judgment with feedback |
| - RuleMaintenance: dynamic confidence tracking + periodic optimization |
| """ |
|
|
| import os |
| import json |
| import re |
| import time |
| import math |
| import hashlib |
| import threading |
| from collections import defaultdict |
| from typing import Dict, List, Tuple, Set, Optional, Any |
| |
| from evorm_mlp_gate import MLPGate |
| |
| from evorm_mlight import MlightRationaleElicitor, MheavyDecisionMaker |
| from evorm_entity_embedding import EntityEmbedder |
| from evorm_config import EvoRMConfig |
| from dataclasses import dataclass, field |
|
|
|
|
| |
| |
| |
|
|
| @dataclass |
| class ConditionAtom: |
| """A single condition atom in a FOL rule.""" |
| atom_type: str |
| attr: str |
| value1: str = "" |
| value2: str = "" |
|
|
| def to_key(self) -> str: |
| return f"{self.atom_type}({self.attr})" |
|
|
| def __hash__(self): |
| return hash(self.to_key()) |
|
|
| def __eq__(self, other): |
| return self.to_key() == other.to_key() |
|
|
|
|
| @dataclass |
| class FOLRule: |
| """A first-order logic rule distilled from LLM inference.""" |
| rule_id: str |
| atoms: List[ConditionAtom] |
| conclusion: int |
| conf: float = 0.0 |
| conf0: float = 0.0 |
| fresh: float = 1.0 |
| trigger_count: int = 0 |
| used_count: int = 0 |
| sR_sum: float = 0.0 |
| last_used_time: float = 0.0 |
| created_time: float = 0.0 |
|
|
| def update_confidence(self): |
| """Update global confidence: Conf(R) = (Conf0 + sum(sR)) / (1 + trigger_count)""" |
| self.conf = (self.conf0 + self.sR_sum) / (1.0 + self.trigger_count) |
|
|
| def update_freshness(self, decay_lambda: float = 0.01): |
| """Update freshness: Fresh(R) = exp(-lambda * (t_now - t_last))""" |
| if self.last_used_time > 0: |
| self.fresh = math.exp(-decay_lambda * (time.time() - self.last_used_time)) |
| else: |
| self.fresh = 1.0 |
|
|
| def atom_set(self) -> Set[str]: |
| return {a.to_key() for a in self.atoms} |
|
|
| def __hash__(self): |
| return hash(self.rule_id) |
|
|
| def __eq__(self, other): |
| return self.rule_id == other.rule_id |
|
|
|
|
| @dataclass |
| class Hyperedge: |
| """A hyperedge in the weighted hypergraph, representing a matching cluster.""" |
| hyperedge_id: str |
| entity_pairs: Set[Tuple[int, int]] = field(default_factory=set) |
| node_set: Set[str] = field(default_factory=set) |
| rules: List[FOLRule] = field(default_factory=list) |
| weight: float = 0.0 |
|
|
| def recompute_weight(self, alpha: float = 0.6, beta: float = 0.4, |
| entity_embeddings: Dict = None): |
| """ |
| Recompute hyperedge weight: wk = alpha * max(Conf(R)) + beta * avg_cos_sim(Vk) |
| |
| Paper Eq. 6: The hyperedge weight combines symbolic rule confidence |
| with neural embedding similarity of entity pairs in the cluster. |
| """ |
| symbolic = 0.0 |
| if self.rules: |
| symbolic = max(r.conf for r in self.rules) |
|
|
| neural = 0.0 |
| if entity_embeddings and self.entity_pairs: |
| |
| embs = [] |
| for ep in self.entity_pairs: |
| eid1, eid2 = f"e_{ep[0]}", f"e_{ep[1]}" |
| if eid1 in entity_embeddings: |
| embs.append(entity_embeddings[eid1]) |
| if eid2 in entity_embeddings: |
| embs.append(entity_embeddings[eid2]) |
| |
| if len(embs) >= 2: |
| |
| from evorm_entity_embedding import EntityEmbedder |
| embedder = EntityEmbedder() |
| neural = embedder.avg_cosine_similarity(embs) |
| elif entity_embeddings: |
| |
| neural = symbolic * 0.5 |
|
|
| self.weight = alpha * symbolic + beta * neural |
|
|
|
|
| |
| |
| |
|
|
| class RuleEncoding: |
| """Convert LLM matching decisions into explicit FOL rules.""" |
|
|
| |
| ATTR_SAME_PATTERN = re.compile( |
| r'(\w+)\s*=\s*same\s*(?:\(([^)]*)\))?', re.IGNORECASE) |
| ATTR_DIFFER_PATTERN = re.compile( |
| r'(\w+)\s*=\s*differ\s*(?:\(([^)]*)\))?', re.IGNORECASE) |
| NEIGHBOR_PATTERN = re.compile( |
| r'neighbor\s*:\s*(\w+)', re.IGNORECASE) |
| NEIGHBOR_DIFFER_PATTERN = re.compile( |
| r'neighbor\s*:\s*differ\s*(\w*)', re.IGNORECASE) |
| SHARE_NEIGHBOR_PATTERN = re.compile( |
| r'(?:share|has)\s+(?:common\s+)?(?:neighbor|relation)\s*(?::\s*)?(\w*)', re.IGNORECASE) |
|
|
| def __init__(self, client=None): |
| self.client = client |
| self.rule_counter = 0 |
|
|
| def _generate_rule_id(self, atoms: List[ConditionAtom]) -> str: |
| """Generate a deterministic rule ID from atoms.""" |
| key = "|".join(sorted(a.to_key() for a in atoms)) |
| return hashlib.md5(key.encode()).hexdigest()[:16] |
|
|
| def parse_rationale(self, rationale: str, conclusion: int) -> List[ConditionAtom]: |
| """ |
| Parse a natural language rationale into typed condition atoms. |
| |
| Args: |
| rationale: natural language rationale from LLM |
| conclusion: 1 (match) or 0 (non-match) |
| |
| Returns: |
| List of ConditionAtom objects |
| """ |
| atoms = [] |
|
|
| |
| decisive_section = "" |
| supporting_section = "" |
|
|
| decisive_match = re.search( |
| r'\[DECISIVE\](.*?)(?:\[SUPPORTING\]|$)', rationale, re.DOTALL) |
| if decisive_match: |
| decisive_section = decisive_match.group(1) |
|
|
| supporting_match = re.search( |
| r'\[SUPPORTING\](.*?)$', rationale, re.DOTALL) |
| if supporting_match: |
| supporting_section = supporting_match.group(1) |
|
|
| all_text = decisive_section + " " + supporting_section |
|
|
| |
| for m in self.ATTR_SAME_PATTERN.finditer(all_text): |
| attr = m.group(1).strip().lower() |
| atoms.append(ConditionAtom( |
| atom_type="SameValue", |
| attr=attr |
| )) |
|
|
| |
| for m in self.ATTR_DIFFER_PATTERN.finditer(all_text): |
| attr = m.group(1).strip().lower() |
| atoms.append(ConditionAtom( |
| atom_type="DifferValue", |
| attr=attr |
| )) |
|
|
| |
| for m in self.NEIGHBOR_PATTERN.finditer(all_text): |
| rel = m.group(1).strip().lower() |
| if rel and 'differ' not in rel.lower(): |
| atoms.append(ConditionAtom( |
| atom_type="ShareNeighbor", |
| attr=rel |
| )) |
|
|
| |
| for m in self.NEIGHBOR_DIFFER_PATTERN.finditer(all_text): |
| rel = m.group(1).strip().lower() |
| if rel: |
| atoms.append(ConditionAtom( |
| atom_type="DifferNeighbor", |
| attr=rel |
| )) |
|
|
| |
| if not atoms: |
| |
| words = re.findall(r'\b\w+\b', all_text) |
| |
| if conclusion == 1: |
| atoms.append(ConditionAtom( |
| atom_type="SameValue", |
| attr="entity_name" |
| )) |
| else: |
| atoms.append(ConditionAtom( |
| atom_type="DifferValue", |
| attr="entity_name" |
| )) |
|
|
| return atoms |
|
|
| def create_rule(self, rationale: str, conclusion: int, |
| used_rules: List[Tuple['FOLRule', float]] = None) -> FOLRule: |
| """ |
| Create a new FOL rule from rationale and LLM feedback. |
| |
| Args: |
| rationale: natural language rationale |
| conclusion: 1 (match) or 0 (non-match) |
| used_rules: list of (rule, contribution_score) from LLM feedback |
| |
| Returns: |
| New FOLRule object |
| """ |
| atoms = self.parse_rationale(rationale, conclusion) |
| rule_id = self._generate_rule_id(atoms) |
|
|
| rule = FOLRule( |
| rule_id=rule_id, |
| atoms=atoms, |
| conclusion=conclusion, |
| created_time=time.time(), |
| last_used_time=time.time(), |
| ) |
|
|
| |
| if used_rules: |
| max_s = 0.0 |
| for used_rule, sR in used_rules: |
| overlap = len(set(a.to_key() for a in atoms) & |
| set(a.to_key() for a in used_rule.atoms)) |
| union = len(set(a.to_key() for a in atoms) | |
| set(a.to_key() for a in used_rule.atoms)) |
| if union > 0 and overlap / union > 0: |
| max_s = max(max_s, sR) |
| rule.conf0 = max_s |
| else: |
| rule.conf0 = 0.0 |
|
|
| rule.conf = rule.conf0 |
| return rule |
|
|
| def elicit_rationale(self, entity_context: str, decision: int, |
| prompt_template: str = None) -> str: |
| """ |
| Use LLM to elicit rationale for a matching decision. |
| |
| Args: |
| entity_context: serialized entity pair context |
| decision: 1 (match) or 0 (non-match) |
| prompt_template: optional custom template |
| |
| Returns: |
| Natural language rationale |
| """ |
| if prompt_template is None: |
| decision_str = "MATCH" if decision == 1 else "NON-MATCH" |
| prompt_template = f"""Analyze the following entity pair and explain why they are a {decision_str}. |
| |
| Entity Context: |
| {entity_context} |
| |
| Please output your analysis in the following format: |
| [DECISIVE] List the core attributes or relations that are pivotal to the {decision_str} decision. |
| [SUPPORTING] List any auxiliary evidence that corroborates the decision. |
| |
| Use the format: attr=same(value), attr=differ(value1 vs value2), or neighbor:rel_name.""" |
| |
| if self.client: |
| try: |
| response = self.client.chat.completions.create( |
| model="gpt-3.5-turbo", |
| messages=[{'role': 'user', 'content': prompt_template}], |
| temperature=0.1 |
| ) |
| return response.choices[0].message.content.strip() |
| except Exception as e: |
| print(f"RuleEncoding: rationale elicitation failed: {e}") |
| return f"[DECISIVE] entity_name=same\n[SUPPORTING] automatic fallback" |
|
|
| |
| return f"[DECISIVE] entity_name={'same' if decision == 1 else 'differ'}\n[SUPPORTING] automatic fallback" |
|
|
|
|
| |
| |
| |
|
|
| class HypergraphStorage: |
| """Weighted hypergraph for organizing rules and entity contexts.""" |
|
|
| def __init__(self, alpha: float = 0.6, beta: float = 0.4, |
| merge_threshold: float = 0.5, |
| entity_embedder: 'EntityEmbedder' = None): |
| self.hyperedges: Dict[str, Hyperedge] = {} |
| self.rules: Dict[str, FOLRule] = {} |
| self.inverted_index: Dict[str, Set[str]] = defaultdict(set) |
| self.alpha = alpha |
| self.beta = beta |
| self.merge_threshold = merge_threshold |
| self.edge_counter = 0 |
| self.entity_embedder = entity_embedder |
| self.entity_embeddings: Dict[str, 'np.ndarray'] = {} |
| self.flat_mode = False |
|
|
| def _generate_edge_id(self) -> str: |
| self.edge_counter += 1 |
| return f"HE_{self.edge_counter:06d}" |
|
|
| def get_or_create_hyperedge(self, entity_pair: Tuple[int, int], |
| node_set: Set[str], |
| rule: FOLRule) -> Hyperedge: |
| """ |
| Find the most similar existing hyperedge or create a new one. |
| |
| Args: |
| entity_pair: (es, et) tuple |
| node_set: VR - nodes extracted from entity context |
| rule: the new FOL rule |
| |
| Returns: |
| The matched or newly created hyperedge |
| """ |
| |
| if self.flat_mode: |
| |
| if rule.rule_id not in self.rules: |
| self.rules[rule.rule_id] = rule |
| |
| dummy = Hyperedge( |
| hyperedge_id='flat', |
| entity_pairs={entity_pair}, |
| node_set=node_set, |
| rules=[rule], |
| ) |
| return dummy |
|
|
| |
| best_id = None |
| best_overlap = 0.0 |
|
|
| for he_id, he in self.hyperedges.items(): |
| if not he.node_set: |
| continue |
| overlap = len(node_set & he.node_set) |
| union = len(node_set | he.node_set) |
| if union > 0: |
| jaccard = overlap / union |
| if jaccard > best_overlap: |
| best_overlap = jaccard |
| best_id = he_id |
|
|
| if best_id and best_overlap >= self.merge_threshold: |
| |
| he = self.hyperedges[best_id] |
| he.entity_pairs.add(entity_pair) |
| he.node_set.update(node_set) |
| if rule not in he.rules: |
| he.rules.append(rule) |
| he.recompute_weight(self.alpha, self.beta, self.entity_embeddings) |
|
|
| |
| for node in node_set: |
| self.inverted_index[node].add(best_id) |
|
|
| return he |
| else: |
| |
| new_id = self._generate_edge_id() |
| he = Hyperedge( |
| hyperedge_id=new_id, |
| entity_pairs={entity_pair}, |
| node_set=node_set, |
| rules=[rule], |
| ) |
| he.recompute_weight(self.alpha, self.beta, self.entity_embeddings) |
| self.hyperedges[new_id] = he |
|
|
| |
| for node in node_set: |
| self.inverted_index[node].add(new_id) |
|
|
| |
| if rule.rule_id not in self.rules: |
| self.rules[rule.rule_id] = rule |
|
|
| return he |
|
|
| def get_candidate_hyperedges(self, entity_ids: List[str]) -> List[Hyperedge]: |
| """ |
| Retrieve candidate hyperedges containing any of the given entity IDs. |
| O(1) complexity via inverted index. |
| |
| Args: |
| entity_ids: list of entity ID strings |
| |
| Returns: |
| List of candidate hyperedges |
| """ |
| candidate_ids: Set[str] = set() |
| for eid in entity_ids: |
| candidate_ids.update(self.inverted_index.get(eid, set())) |
|
|
| return [self.hyperedges[hid] for hid in candidate_ids |
| if hid in self.hyperedges] |
|
|
| def store_entity_embedding(self, entity_id: str, embedding: 'np.ndarray'): |
| """Store entity embedding for neural similarity computation.""" |
| self.entity_embeddings[entity_id] = embedding |
|
|
| def get_entity_embeddings(self) -> Dict: |
| """Get all stored entity embeddings.""" |
| return self.entity_embeddings |
|
|
| def get_candidate_rules(self, entity_ids: List[str]) -> List[FOLRule]: |
| """ |
| Get all candidate rules from hyperedges containing the entity IDs. |
| |
| Args: |
| entity_ids: list of entity ID strings |
| |
| Returns: |
| List of candidate FOL rules |
| """ |
| |
| if self.flat_mode: |
| return list(self.rules.values()) |
|
|
| candidates = self.get_candidate_hyperedges(entity_ids) |
| rules_set: Dict[str, FOLRule] = {} |
| for he in candidates: |
| for rule in he.rules: |
| if rule.rule_id not in rules_set: |
| rules_set[rule.rule_id] = rule |
| return list(rules_set.values()) |
|
|
| def register_rule(self, rule: FOLRule): |
| """Register a rule in the global rule set.""" |
| if rule.rule_id not in self.rules: |
| self.rules[rule.rule_id] = rule |
|
|
| def get_rule(self, rule_id: str) -> Optional[FOLRule]: |
| return self.rules.get(rule_id) |
|
|
| def stats(self) -> Dict: |
| return { |
| 'num_hyperedges': len(self.hyperedges), |
| 'num_rules': len(self.rules), |
| 'num_indexed_entities': len(self.inverted_index), |
| 'avg_rules_per_edge': ( |
| sum(len(he.rules) for he in self.hyperedges.values()) / |
| max(1, len(self.hyperedges)) |
| ), |
| } |
|
|
|
|
| |
| |
| |
|
|
| class TwoStageInferenceController: |
| """Two-stage controller: symbolic filtering + LLM judgment with feedback.""" |
|
|
| def __init__(self, hypergraph: HypergraphStorage, |
| client=None, |
| theta_hi: float = 0.7, |
| theta_prune: float = 0.5, |
| theta_gate: float = 0.3, |
| K: int = 5, |
| use_mlight: bool = True): |
| self.hypergraph = hypergraph |
| self.client = client |
| self.theta_hi = theta_hi |
| self.theta_prune = theta_prune |
| self.theta_gate = theta_gate |
| self.K = K |
|
|
| |
| self.mlp_gate = None |
| self.use_mlp_gate = True |
|
|
| |
| self.use_mlight = use_mlight |
| self.mlight = None |
| self.mheavy = None |
|
|
| |
| self.stage1_hits = 0 |
| self.stage1_total = 0 |
| self.stage2_hits = 0 |
| self.stage2_total = 0 |
|
|
| def verify_rule(self, rule: FOLRule, es_context: Dict, |
| et_context: Dict) -> bool: |
| """ |
| Verify if a rule is triggered for the given entity pair. |
| |
| Args: |
| rule: FOL rule to verify |
| es_context: source entity context (dict with attributes) |
| et_context: target entity context (dict with attributes) |
| |
| Returns: |
| True if all atoms in the rule hold |
| """ |
| for atom in rule.atoms: |
| if not self._verify_atom(atom, es_context, et_context): |
| return False |
| return True |
|
|
| def _verify_atom(self, atom: ConditionAtom, es_ctx: Dict, |
| et_ctx: Dict) -> bool: |
| """Verify a single condition atom.""" |
| attr = atom.attr |
|
|
| if atom.atom_type == "SameValue": |
| |
| v1 = es_ctx.get(attr, "") |
| v2 = et_ctx.get(attr, "") |
| return v1 != "" and v2 != "" and v1.lower() == v2.lower() |
|
|
| elif atom.atom_type == "DifferValue": |
| v1 = es_ctx.get(attr, "") |
| v2 = et_ctx.get(attr, "") |
| return v1 != "" and v2 != "" and v1.lower() != v2.lower() |
|
|
| elif atom.atom_type == "ShareNeighbor": |
| |
| es_neighbors = es_ctx.get(f"neighbors_{attr}", set()) |
| et_neighbors = et_ctx.get(f"neighbors_{attr}", set()) |
| return bool(es_neighbors & et_neighbors) |
|
|
| elif atom.atom_type == "DifferNeighbor": |
| es_neighbors = es_ctx.get(f"neighbors_{attr}", set()) |
| et_neighbors = et_ctx.get(f"neighbors_{attr}", set()) |
| return not bool(es_neighbors & et_neighbors) and es_neighbors and et_neighbors |
|
|
| elif atom.atom_type == "SemanticEquiv": |
| |
| v1 = es_ctx.get(attr, "") |
| v2 = et_ctx.get(attr, "") |
| if v1 and v2: |
| |
| tokens1 = set(v1.lower().split()) |
| tokens2 = set(v2.lower().split()) |
| if tokens1 and tokens2: |
| overlap = len(tokens1 & tokens2) / len(tokens1 | tokens2) |
| return overlap > 0.5 |
| return False |
|
|
| elif atom.atom_type == "SemanticConflict": |
| v1 = es_ctx.get(attr, "") |
| v2 = et_ctx.get(attr, "") |
| if v1 and v2: |
| tokens1 = set(v1.lower().split()) |
| tokens2 = set(v2.lower().split()) |
| if tokens1 and tokens2: |
| overlap = len(tokens1 & tokens2) / len(tokens1 | tokens2) |
| return overlap < 0.2 |
| return False |
|
|
| return False |
|
|
| def stage1_symbolic_filtering(self, es_id: str, et_id: str, |
| es_context: Dict, et_context: Dict |
| ) -> Tuple[Optional[int], List[FOLRule], List[FOLRule]]: |
| """ |
| Stage 1: Symbolic filtering and gated routing. |
| |
| Returns: |
| (decision, triggered_rules, candidate_rules) |
| decision: None if survival (needs Stage 2), 0 or 1 if direct routing |
| """ |
| self.stage1_total += 1 |
|
|
| |
| candidate_rules = self.hypergraph.get_candidate_rules([es_id, et_id]) |
|
|
| |
| triggered: List[FOLRule] = [] |
| for rule in candidate_rules: |
| if self.verify_rule(rule, es_context, et_context): |
| triggered.append(rule) |
|
|
| if not triggered: |
| return None, [], candidate_rules |
|
|
| |
| match_rules = [r for r in triggered if r.conclusion == 1] |
| nonmatch_rules = [r for r in triggered if r.conclusion == 0] |
|
|
| |
| if match_rules and not nonmatch_rules: |
| max_conf = max(r.conf for r in match_rules) |
| if max_conf >= self.theta_hi: |
| self.stage1_hits += 1 |
| return 1, triggered, candidate_rules |
|
|
| |
| if nonmatch_rules and not match_rules: |
| max_conf = max(r.conf for r in nonmatch_rules) |
| if max_conf >= self.theta_prune: |
| self.stage1_hits += 1 |
| return 0, triggered, candidate_rules |
|
|
| |
| if self.use_mlp_gate and self.mlp_gate is not None and self.mlp_gate.is_trained: |
| try: |
| features = self.mlp_gate.extract_features(es_context, et_context, triggered) |
| if not self.mlp_gate.should_invoke_llm(features): |
| self.stage1_hits += 1 |
| return 0, triggered, candidate_rules |
| except Exception: |
| pass |
| |
| return None, triggered, candidate_rules |
|
|
| def construct_evidence_subgraph(self, es_id: str, et_id: str, |
| candidate_rules: List[FOLRule]) -> Dict: |
| """ |
| Construct a compact evidence subgraph for LLM inference. |
| |
| Args: |
| es_id: source entity ID |
| et_id: target entity ID |
| candidate_rules: candidate rules |
| |
| Returns: |
| Evidence subgraph data |
| """ |
| |
| candidate_edges = self.hypergraph.get_candidate_hyperedges([es_id, et_id]) |
|
|
| |
| joint_edges = [] |
| other_edges = [] |
| for he in candidate_edges: |
| if es_id in he.node_set and et_id in he.node_set: |
| joint_edges.append(he) |
| else: |
| other_edges.append(he) |
|
|
| |
| joint_edges.sort(key=lambda x: x.weight, reverse=True) |
| other_edges.sort(key=lambda x: x.weight, reverse=True) |
|
|
| |
| selected = joint_edges[:self.K] |
| if len(selected) < self.K: |
| selected.extend(other_edges[:self.K - len(selected)]) |
|
|
| return { |
| 'hyperedges': selected, |
| 'rules': candidate_rules, |
| } |
|
|
| def build_stage2_prompt(self, es_id: str, et_id: str, |
| es_context: Dict, et_context: Dict, |
| evidence: Dict, |
| triggered_rules: List[FOLRule], |
| base_prompt: str) -> str: |
| """ |
| Build the Stage 2 prompt with evidence subgraph and rule metadata. |
| |
| Args: |
| es_id, et_id: entity IDs |
| es_context, et_context: entity contexts |
| evidence: evidence subgraph from construct_evidence_subgraph |
| triggered_rules: triggered rules |
| base_prompt: original prompt from AdaCoAgentEA |
| |
| Returns: |
| Augmented prompt |
| """ |
| |
| rule_section = "" |
| if triggered_rules: |
| rule_section = "\n\n**Triggered Rules (from historical matching experience):**\n" |
| for i, rule in enumerate(triggered_rules, 1): |
| atom_strs = [] |
| for a in rule.atoms: |
| atom_strs.append(f"{a.atom_type}({a.attr})") |
| atoms_str = " ∧ ".join(atom_strs) |
| verdict = "MATCH" if rule.conclusion == 1 else "NON-MATCH" |
| rule_section += ( |
| f" Rule {i}: {atoms_str} → {verdict}\n" |
| f" Confidence: {rule.conf:.3f} | " |
| f"Triggered: {rule.trigger_count} times\n" |
| ) |
|
|
| |
| evidence_section = "" |
| if evidence.get('hyperedges'): |
| evidence_section = "\n\n**Evidence from Similar Historical Cases:**\n" |
| for he in evidence['hyperedges'][:3]: |
| matching_pairs = len(he.entity_pairs) |
| evidence_section += ( |
| f" - Cluster (weight={he.weight:.3f}): " |
| f"{matching_pairs} similar entity pairs stored\n" |
| ) |
|
|
| |
| feedback_instruction = """ |
| |
| **IMPORTANT: Rule Contribution Feedback** |
| After making your decision, you MUST also provide a JSON object with contribution scores (0.0 to 1.0) for each triggered rule that influenced your decision. Format: |
| ```json |
| { |
| "decision": "MATCH" or "NON-MATCH", |
| "rule_feedback": { |
| "rule_1": 0.8, |
| "rule_2": 0.3 |
| }, |
| "reasoning": "[DECISIVE] ... [SUPPORTING] ..." |
| } |
| ``` |
| Score meaning: 1.0 = DECISIVE (rule was the key factor), 0.1-0.9 = SUPPORTING (rule partially influenced), 0.0 = IRRELEVANT (rule was considered but not used). |
| |
| You MUST respond with valid JSON only.""" |
|
|
| return base_prompt + rule_section + evidence_section + feedback_instruction |
|
|
| def parse_llm_feedback(self, response_text: str |
| ) -> Tuple[Optional[int], Dict[str, float], str]: |
| """ |
| Parse the LLM's JSON response to extract decision, rule feedback, and rationale. |
| |
| Returns: |
| (decision, rule_feedback dict, rationale) |
| """ |
| |
| try: |
| |
| json_match = re.search(r'```json\s*(.*?)\s*```', response_text, re.DOTALL) |
| if json_match: |
| json_str = json_match.group(1) |
| else: |
| |
| json_match = re.search(r'\{.*"decision".*\}', response_text, re.DOTALL) |
| if json_match: |
| json_str = json_match.group(0) |
| else: |
| json_str = response_text |
|
|
| data = json.loads(json_str) |
| decision_str = data.get('decision', '').upper() |
| decision = 1 if 'MATCH' in decision_str and 'NON' not in decision_str else 0 |
|
|
| rule_feedback = data.get('rule_feedback', {}) |
| reasoning = data.get('reasoning', '') |
|
|
| return decision, rule_feedback, reasoning |
| except (json.JSONDecodeError, KeyError): |
| pass |
|
|
| |
| decision = None |
| if 'match' in response_text.lower() and 'non-match' not in response_text.lower(): |
| decision = 1 |
| elif 'non-match' in response_text.lower() or 'no match' in response_text.lower(): |
| decision = 0 |
|
|
| return decision, {}, response_text |
|
|
| def stage2_elicit_rationale(self, es_id: str, et_id: str, |
| es_context: Dict, et_context: Dict, |
| triggered_rules: List[FOLRule]): |
| """ |
| Stage 2a: Elicit rationale using Mlight (independent from decision). |
| |
| Paper Section III-C Step 1: Mlight generates a natural language |
| rationale explaining the matching evidence, WITHOUT making a decision. |
| |
| Returns: |
| RationaleResult with structured analysis |
| """ |
| if self.mlight is not None: |
| return self.mlight.elicit(es_context, et_context, triggered_rules) |
| else: |
| from evorm_mlight import RationaleResult |
| return RationaleResult( |
| rationale="[DECISIVE] entity_name=same\\n[SUPPORTING] automatic fallback (no Mlight)", |
| token_usage=0, |
| ) |
|
|
| def stage2_llm_judgment(self, es_id: str, et_id: str, |
| es_context: Dict, et_context: Dict, |
| triggered_rules: List[FOLRule], |
| candidate_rules: List[FOLRule], |
| base_prompt: str) -> Dict: |
| """ |
| Stage 2: LLM judgment with rule feedback. |
| |
| Returns: |
| Dict with decision, rule_feedback, rationale, token_usage |
| """ |
| self.stage2_total += 1 |
|
|
| |
| for rule in triggered_rules: |
| rule.trigger_count += 1 |
|
|
| |
| |
| |
| if self.use_mlight and self.mlight is not None and self.mheavy is not None: |
| |
| rationale_result = self.stage2_elicit_rationale( |
| es_id, et_id, es_context, et_context, triggered_rules) |
|
|
| |
| decision_result = self.mheavy.decide( |
| es_context, et_context, |
| rationale_result.rationale, |
| triggered_rules) |
|
|
| total_tokens = (rationale_result.token_usage + |
| decision_result.get('token_usage', 0)) |
|
|
| if decision_result.get('decision') is not None: |
| self.stage2_hits += 1 |
|
|
| return { |
| 'decision': decision_result.get('decision'), |
| 'rule_feedback': decision_result.get('rule_feedback', {}), |
| 'rationale': rationale_result.rationale, |
| 'confidence': decision_result.get('confidence', 0.5), |
| 'token_usage': total_tokens, |
| 'raw_response': decision_result.get('raw_response', ''), |
| 'mlight_raw': rationale_result.raw_response, |
| 'method': 'mlight+mheavy', |
| } |
|
|
| |
| |
| |
| |
| evidence = self.construct_evidence_subgraph( |
| es_id, et_id, candidate_rules) |
|
|
| |
| prompt = self.build_stage2_prompt( |
| es_id, et_id, es_context, et_context, |
| evidence, triggered_rules, base_prompt) |
|
|
| if self.client: |
| try: |
| response = self.client.chat.completions.create( |
| model="gpt-3.5-turbo", |
| messages=[{'role': 'user', 'content': prompt}], |
| temperature=0.1, |
| response_format={"type": "json_object"}, |
| ) |
| response_text = response.choices[0].message.content.strip() |
| tokens = response.usage.total_tokens |
|
|
| decision, rule_feedback, rationale = self.parse_llm_feedback( |
| response_text) |
|
|
| if decision is not None: |
| self.stage2_hits += 1 |
|
|
| return { |
| 'decision': decision, |
| 'rule_feedback': rule_feedback, |
| 'rationale': rationale, |
| 'token_usage': tokens, |
| 'raw_response': response_text, |
| 'method': 'single-call', |
| } |
| except Exception as e: |
| print(f"Stage 2 LLM call failed: {e}") |
| return { |
| 'decision': None, |
| 'rule_feedback': {}, |
| 'rationale': '', |
| 'token_usage': 0, |
| 'raw_response': '', |
| 'error': str(e), |
| 'method': 'error', |
| } |
|
|
| |
| return { |
| 'decision': None, |
| 'rule_feedback': {}, |
| 'rationale': 'no client available', |
| 'token_usage': 0, |
| 'raw_response': '', |
| 'method': 'no-client', |
| } |
|
|
| def get_stats(self) -> Dict: |
| stats = { |
| 'stage1_total': self.stage1_total, |
| 'stage1_hits': self.stage1_hits, |
| 'stage1_rate': self.stage1_hits / max(1, self.stage1_total), |
| 'stage2_total': self.stage2_total, |
| 'stage2_hits': self.stage2_hits, |
| 'use_mlight': self.use_mlight, |
| } |
| if self.mlight is not None: |
| stats['mlight'] = self.mlight.get_stats() |
| if self.mheavy is not None: |
| stats['mheavy'] = self.mheavy.get_stats() |
| return stats |
|
|
|
|
| |
| |
| |
|
|
| class RuleMaintenance: |
| """Dynamic rule maintenance: confidence tracking + periodic optimization.""" |
|
|
| def __init__(self, hypergraph: HypergraphStorage, |
| decay_lambda: float = 0.01, |
| freshness_threshold: float = 0.1, |
| confidence_threshold: float = 0.2, |
| eval_triggers: int = 10, |
| merge_similarity: float = 0.7, |
| max_rules: int = 10000, |
| evict_percentile: float = 0.1): |
| self.hypergraph = hypergraph |
| self.decay_lambda = decay_lambda |
| self.theta_f = freshness_threshold |
| self.theta_c = confidence_threshold |
| self.N_eval = eval_triggers |
| self.eta_m = merge_similarity |
| self.max_rules = max_rules |
| self.evict_percentile = evict_percentile |
|
|
| self.last_optimization_time = time.time() |
| self.optimization_interval = 300 |
|
|
| def update_rule_confidence(self, rule_id: str, sR: float): |
| """ |
| Update rule confidence with LLM contribution score. |
| |
| Args: |
| rule_id: the rule to update |
| sR: contribution score from LLM (0.0 to 1.0) |
| """ |
| rule = self.hypergraph.get_rule(rule_id) |
| if rule is None: |
| return |
|
|
| rule.sR_sum += sR |
| if sR > 0: |
| rule.used_count += 1 |
| rule.last_used_time = time.time() |
|
|
| rule.update_confidence() |
| rule.update_freshness(self.decay_lambda) |
|
|
| def should_optimize(self) -> bool: |
| """Check if it's time for periodic optimization.""" |
| return (time.time() - self.last_optimization_time) >= self.optimization_interval |
|
|
| def optimize(self): |
| """Execute periodic set optimization: pruning, flipping, merging, forgetting.""" |
| self.last_optimization_time = time.time() |
|
|
| rules = list(self.hypergraph.rules.values()) |
| if not rules: |
| return |
|
|
| |
| stale_rules = [] |
| for rule in rules: |
| rule.update_freshness(self.decay_lambda) |
| if rule.fresh < self.theta_f: |
| stale_rules.append(rule.rule_id) |
|
|
| for rid in stale_rules: |
| self._remove_rule(rid) |
| if stale_rules: |
| print(f"RuleMaintenance: pruned {len(stale_rules)} stale rules") |
|
|
| |
| remaining = [r for r in self.hypergraph.rules.values() |
| if r.rule_id not in stale_rules] |
| for rule in remaining: |
| if (rule.conf < self.theta_c and |
| rule.trigger_count >= self.N_eval and |
| rule.used_count < rule.trigger_count * 0.3): |
| |
| old_conclusion = rule.conclusion |
| rule.conclusion = 1 - rule.conclusion |
| rule.conf = 0.3 |
| rule.conf0 = 0.3 |
| rule.sR_sum = 0.0 |
| rule.used_count = 0 |
| print(f"RuleMaintenance: flipped conclusion of {rule.rule_id} " |
| f"from {old_conclusion} to {rule.conclusion}") |
|
|
| |
| self._merge_similar_rules() |
|
|
| |
| if len(self.hypergraph.rules) > self.max_rules: |
| self._evict_low_utility_rules() |
|
|
| def _remove_rule(self, rule_id: str): |
| """Remove a rule from the hypergraph.""" |
| if rule_id in self.hypergraph.rules: |
| del self.hypergraph.rules[rule_id] |
|
|
| |
| for he in self.hypergraph.hyperedges.values(): |
| he.rules = [r for r in he.rules if r.rule_id != rule_id] |
|
|
| def _merge_similar_rules(self): |
| """Merge similar rules using Jaccard similarity on atom sets.""" |
| rule_list = list(self.hypergraph.rules.values()) |
| merged = set() |
|
|
| for i, r1 in enumerate(rule_list): |
| if r1.rule_id in merged: |
| continue |
| cluster = [r1] |
| for j, r2 in enumerate(rule_list): |
| if i >= j or r2.rule_id in merged: |
| continue |
| if r1.conclusion == r2.conclusion: |
| set1 = r1.atom_set() |
| set2 = r2.atom_set() |
| if set1 and set2: |
| union = len(set1 | set2) |
| if union > 0: |
| jaccard = len(set1 & set2) / union |
| if jaccard >= self.eta_m: |
| cluster.append(r2) |
|
|
| if len(cluster) > 1: |
| |
| best = max(cluster, key=lambda r: r.conf * r.trigger_count) |
| for r in cluster: |
| if r.rule_id != best.rule_id: |
| merged.add(r.rule_id) |
| |
| best.trigger_count += r.trigger_count |
| best.sR_sum += r.sR_sum |
| best.used_count += r.used_count |
| best.last_used_time = max( |
| best.last_used_time, r.last_used_time) |
| self._remove_rule(r.rule_id) |
|
|
| |
| total_trig = sum(r.trigger_count for r in cluster) |
| if total_trig > 0: |
| best.conf = sum(r.conf * r.trigger_count |
| for r in cluster) / total_trig |
| best.conf0 = best.conf |
| best.update_freshness(self.decay_lambda) |
| print(f"RuleMaintenance: merged {len(cluster)} rules into {best.rule_id}") |
|
|
| |
| for rid in list(merged): |
| if rid in self.hypergraph.rules: |
| del self.hypergraph.rules[rid] |
|
|
| def _evict_low_utility_rules(self): |
| """Evict low-utility rules when capacity is exceeded.""" |
| rules = list(self.hypergraph.rules.values()) |
| |
| rules.sort(key=lambda r: r.conf * r.fresh) |
|
|
| num_to_evict = int(len(rules) * self.evict_percentile) |
| for rule in rules[:num_to_evict]: |
| self._remove_rule(rule.rule_id) |
|
|
| print(f"RuleMaintenance: evicted {num_to_evict} low-utility rules") |
|
|
|
|
| |
| |
| |
|
|
| class EvoRMPlugin: |
| """ |
| Main EvoRM plugin class for integration with AdaCoAgentEA. |
| |
| Usage: |
| plugin = EvoRMPlugin(client=openai_client) |
| |
| # Stage 1: Check if symbolic routing can decide |
| decision, triggered, candidates = plugin.stage1(es_id, et_id, es_ctx, et_ctx) |
| |
| if decision is not None: |
| # Direct routing - no LLM needed |
| return decision |
| else: |
| # Stage 2: LLM with evidence and rule feedback |
| result = plugin.stage2(es_id, et_id, es_ctx, et_ctx, |
| triggered, candidates, base_prompt) |
| # Process result['decision'] and result['rule_feedback'] |
| """ |
|
|
| def __init__(self, client=None, |
| theta_hi: float = 0.7, |
| theta_prune: float = 0.5, |
| alpha: float = 0.6, |
| beta: float = 0.4, |
| merge_threshold: float = 0.5, |
| persistence_dir: str = None, |
| ablation_mode: str = None, |
| config: 'EvoRMConfig' = None): |
| |
| if config is not None: |
| self.config = config |
| alpha = config.alpha |
| beta = config.beta |
| merge_threshold = config.merge_threshold |
| theta_hi = config.theta_hi |
| theta_prune = config.theta_prune |
| persistence_dir = config.persistence_dir or persistence_dir |
| ablation_mode = config.ablation_mode or ablation_mode |
| else: |
| self.config = EvoRMConfig( |
| alpha=alpha, beta=beta, merge_threshold=merge_threshold, |
| theta_hi=theta_hi, theta_prune=theta_prune, |
| persistence_dir=persistence_dir, ablation_mode=ablation_mode) |
| |
| self.client = client |
| self.persistence_dir = persistence_dir |
| self.ablation_mode = ablation_mode |
|
|
| |
| self.rule_encoding = RuleEncoding(client=client) |
| self.hypergraph = HypergraphStorage( |
| alpha=alpha, beta=beta, merge_threshold=merge_threshold) |
| self.controller = TwoStageInferenceController( |
| self.hypergraph, client=client, |
| theta_hi=theta_hi, theta_prune=theta_prune) |
| self.maintenance = RuleMaintenance( |
| self.hypergraph) |
|
|
| |
| self.entity_embedder = EntityEmbedder( |
| demb=self.config.demb, n_features=self.config.n_features) |
| self.hypergraph.entity_embedder = self.entity_embedder |
|
|
| |
| self.mlight = MlightRationaleElicitor( |
| client=client, model=self.config.model, |
| temperature=self.config.mlight_temperature, |
| max_retries=self.config.api_max_retries, |
| timeout=self.config.api_timeout) |
| self.mheavy = MheavyDecisionMaker( |
| client=client, model=self.config.model, |
| temperature=self.config.mheavy_temperature, |
| max_retries=self.config.api_max_retries, |
| timeout=self.config.api_timeout) |
| self.controller.mlight = self.mlight |
| self.controller.mheavy = self.mheavy |
|
|
| |
| if self.ablation_mode == 'no_mlight': |
| self.controller.use_mlight = False |
|
|
| |
| if self.ablation_mode == 'no_hypergraph': |
| self.hypergraph.flat_mode = True |
|
|
| |
| try: |
| import torch |
| device = 'cuda' if torch.cuda.is_available() else 'cpu' |
| except ImportError: |
| device = 'cpu' |
| self.mlp_gate = MLPGate( |
| input_dim=self.config.mlp_input_dim, |
| hidden_dims=self.config.mlp_hidden_dims, |
| theta_gate=self.config.theta_gate, |
| n_warmup=self.config.mlp_n_warmup, |
| device=device, |
| ) |
| self.controller.mlp_gate = self.mlp_gate |
| self._mlp_samples_collected = 0 |
|
|
| |
| if self.ablation_mode == 'no_mlp_gate': |
| self.controller.use_mlp_gate = False |
|
|
| |
| self.total_queries = 0 |
| self.llm_calls_saved = 0 |
| self.total_tokens = 0 |
|
|
| |
| if persistence_dir: |
| self._load_state() |
|
|
| def stage1(self, es_id: str, et_id: str, |
| es_context: Dict, et_context: Dict |
| ) -> Tuple[Optional[int], List[FOLRule], List[FOLRule]]: |
| """ |
| Stage 1: Symbolic filtering. |
| |
| Returns: |
| (decision, triggered_rules, candidate_rules) |
| decision: None if needs Stage 2, 0/1 if direct routing |
| """ |
| |
| if self.ablation_mode == 'no_stage1': |
| return None, [], [] |
| self.total_queries += 1 |
| return self.controller.stage1_symbolic_filtering( |
| es_id, et_id, es_context, et_context) |
|
|
| def stage2(self, es_id: str, et_id: str, |
| es_context: Dict, et_context: Dict, |
| triggered_rules: List[FOLRule], |
| candidate_rules: List[FOLRule], |
| base_prompt: str) -> Dict: |
| """ |
| Stage 2: LLM judgment with evidence and rule feedback. |
| |
| Returns: |
| Dict with 'decision', 'rule_feedback', 'rationale', 'token_usage' |
| """ |
| result = self.controller.stage2_llm_judgment( |
| es_id, et_id, es_context, et_context, |
| triggered_rules, candidate_rules, base_prompt) |
|
|
| self.total_tokens += result.get('token_usage', 0) |
|
|
| |
| if result.get('rule_feedback'): |
| for rule_key, sR in result['rule_feedback'].items(): |
| |
| rule_id = self._find_rule_by_key(rule_key, triggered_rules) |
| if rule_id: |
| self.maintenance.update_rule_confidence(rule_id, float(sR)) |
|
|
| return result |
|
|
| def _find_rule_by_key(self, key: str, rules: List[FOLRule]) -> Optional[str]: |
| """Find a rule ID by its index key (e.g., 'rule_1').""" |
| match = re.match(r'rule_(\d+)', key) |
| if match: |
| idx = int(match.group(1)) - 1 |
| if 0 <= idx < len(rules): |
| return rules[idx].rule_id |
| return None |
|
|
| def record_trajectory(self, es_id: str, et_id: str, |
| es_context: Dict, et_context: Dict, |
| decision: int, rationale: str, |
| triggered_rules: List[FOLRule], |
| rule_feedback: Dict[str, float]): |
| """ |
| Record a matching trajectory after LLM inference. |
| Creates/updates hyperedges and rules in the hypergraph. |
| |
| Args: |
| es_id: source entity ID |
| et_id: target entity ID |
| es_context: source entity context |
| et_context: target entity context |
| decision: 1 (match) or 0 (non-match) |
| rationale: LLM rationale |
| triggered_rules: triggered rules |
| rule_feedback: dict of rule_key -> contribution score |
| """ |
| |
| if self.ablation_mode == 'no_maintenance': |
| return |
| |
| used_rules = [] |
| for rule_key, sR in rule_feedback.items(): |
| rule_id = self._find_rule_by_key(rule_key, triggered_rules) |
| if rule_id: |
| rule = self.hypergraph.get_rule(rule_id) |
| if rule: |
| used_rules.append((rule, float(sR))) |
|
|
| new_rule = self.rule_encoding.create_rule( |
| rationale, decision, used_rules) |
|
|
| |
| node_set = set() |
| node_set.add(f"e_{es_id}") |
| node_set.add(f"e_{et_id}") |
|
|
| for key, val in es_context.items(): |
| if not key.startswith('neighbors_'): |
| node_set.add(f"a_{key}") |
| if val: |
| node_set.add(f"v_{key}_{str(val)[:50]}") |
|
|
| for key, val in et_context.items(): |
| if not key.startswith('neighbors_'): |
| node_set.add(f"a_{key}") |
| if val: |
| node_set.add(f"v_{key}_{str(val)[:50]}") |
|
|
| |
| if hasattr(self, 'entity_embedder') and self.entity_embedder is not None: |
| try: |
| emb_s = self.entity_embedder.embed_entity(es_context, f"e_{es_id}") |
| emb_t = self.entity_embedder.embed_entity(et_context, f"e_{et_id}") |
| self.hypergraph.store_entity_embedding(f"e_{es_id}", emb_s) |
| self.hypergraph.store_entity_embedding(f"e_{et_id}", emb_t) |
| except Exception: |
| pass |
|
|
| |
| self.hypergraph.get_or_create_hyperedge( |
| (int(es_id) if es_id.isdigit() else hash(es_id), |
| int(et_id) if et_id.isdigit() else hash(et_id)), |
| node_set, new_rule) |
|
|
| |
| if hasattr(self, 'mlp_gate') and self.mlp_gate is not None: |
| if not self.mlp_gate.is_trained and triggered_rules: |
| try: |
| features = self.mlp_gate.extract_features( |
| es_context, et_context, triggered_rules) |
| self.mlp_gate.collect_sample(features, float(decision)) |
| self._mlp_samples_collected += 1 |
| except Exception: |
| pass |
|
|
| |
| if self.maintenance.should_optimize(): |
| self.maintenance.optimize() |
|
|
| def build_context_dict(self, entity_name: str, relations: List[str] = None, |
| descriptions: str = "") -> Dict: |
| """ |
| Build a context dictionary from entity information. |
| |
| Args: |
| entity_name: entity name |
| relations: list of relation strings (e.g., "Has relation 'X' with Y") |
| descriptions: entity description text |
| |
| Returns: |
| Context dict suitable for rule verification |
| """ |
| ctx = {'entity_name': entity_name, 'description': descriptions} |
|
|
| if relations: |
| for rel_str in relations: |
| |
| m1 = re.match(r"Has relation '([^']+)' with (.+)", rel_str) |
| if m1: |
| rel = m1.group(1).lower() |
| neighbor = m1.group(2).strip() |
| key = f"neighbors_{rel}" |
| if key not in ctx: |
| ctx[key] = set() |
| ctx[key].add(neighbor) |
|
|
| m2 = re.match(r"Is ([^']+) of (.+)", rel_str) |
| if m2: |
| rel = m2.group(1).lower() |
| neighbor = m2.group(2).strip() |
| key = f"neighbors_{rel}" |
| if key not in ctx: |
| ctx[key] = set() |
| ctx[key].add(neighbor) |
|
|
| return ctx |
|
|
| def get_stats(self) -> Dict: |
| """Get comprehensive statistics.""" |
| controller_stats = self.controller.get_stats() |
| hypergraph_stats = self.hypergraph.stats() |
|
|
| stats = { |
| **controller_stats, |
| **hypergraph_stats, |
| 'total_queries': self.total_queries, |
| 'llm_calls_saved': self.llm_calls_saved, |
| 'llm_save_rate': (self.llm_calls_saved / |
| max(1, self.llm_calls_saved + self.controller.stage2_total)), |
| 'total_tokens': self.total_tokens, |
| } |
| if hasattr(self, 'mlp_gate') and self.mlp_gate is not None: |
| stats['mlp_gate'] = self.mlp_gate.get_stats() |
| if hasattr(self, 'entity_embedder') and self.entity_embedder is not None: |
| stats['entity_embedder'] = self.entity_embedder.get_stats() |
| return stats |
|
|
| def _save_state(self): |
| """Persist hypergraph state to disk.""" |
| if not self.persistence_dir: |
| return |
| os.makedirs(self.persistence_dir, exist_ok=True) |
|
|
| state = { |
| 'rules': {}, |
| 'hyperedges': {}, |
| 'inverted_index': {k: list(v) for k, v in self.hypergraph.inverted_index.items()}, |
| 'stats': self.get_stats(), |
| } |
|
|
| for rid, rule in self.hypergraph.rules.items(): |
| state['rules'][rid] = { |
| 'rule_id': rule.rule_id, |
| 'atoms': [{'atom_type': a.atom_type, 'attr': a.attr, |
| 'value1': a.value1, 'value2': a.value2} |
| for a in rule.atoms], |
| 'conclusion': rule.conclusion, |
| 'conf': rule.conf, |
| 'conf0': rule.conf0, |
| 'trigger_count': rule.trigger_count, |
| 'used_count': rule.used_count, |
| 'sR_sum': rule.sR_sum, |
| 'last_used_time': rule.last_used_time, |
| 'created_time': rule.created_time, |
| } |
|
|
| for hid, he in self.hypergraph.hyperedges.items(): |
| state['hyperedges'][hid] = { |
| 'hyperedge_id': he.hyperedge_id, |
| 'entity_pairs': list(he.entity_pairs), |
| 'node_set': list(he.node_set), |
| 'rule_ids': [r.rule_id for r in he.rules], |
| 'weight': he.weight, |
| } |
|
|
| state_path = os.path.join(self.persistence_dir, 'evorm_state.json') |
| with open(state_path, 'w', encoding='utf-8') as f: |
| json.dump(state, f, ensure_ascii=False, indent=2) |
|
|
| |
| if hasattr(self, 'mlp_gate') and self.mlp_gate is not None: |
| try: |
| mlp_path = os.path.join(self.persistence_dir, 'mlp_gate.pt') |
| self.mlp_gate.save(mlp_path) |
| except Exception as e: |
| print(f"Failed to save MLP gate: {e}") |
|
|
| print(f"EvoRM state saved to {state_path}") |
|
|
| def _load_state(self): |
| """Load persisted hypergraph state from disk.""" |
| if not self.persistence_dir: |
| return |
| state_path = os.path.join(self.persistence_dir, 'evorm_state.json') |
| if not os.path.exists(state_path): |
| return |
|
|
| try: |
| with open(state_path, 'r', encoding='utf-8') as f: |
| state = json.load(f) |
|
|
| |
| for rid, rdata in state.get('rules', {}).items(): |
| atoms = [ConditionAtom(**a) for a in rdata['atoms']] |
| rule = FOLRule( |
| rule_id=rdata['rule_id'], |
| atoms=atoms, |
| conclusion=rdata['conclusion'], |
| conf=rdata['conf'], |
| conf0=rdata['conf0'], |
| trigger_count=rdata['trigger_count'], |
| used_count=rdata['used_count'], |
| sR_sum=rdata['sR_sum'], |
| last_used_time=rdata['last_used_time'], |
| created_time=rdata['created_time'], |
| ) |
| self.hypergraph.rules[rid] = rule |
|
|
| |
| for hid, hdata in state.get('hyperedges', {}).items(): |
| he = Hyperedge( |
| hyperedge_id=hdata['hyperedge_id'], |
| entity_pairs=set( |
| tuple(p) for p in hdata['entity_pairs']), |
| node_set=set(hdata['node_set']), |
| weight=hdata['weight'], |
| ) |
| for rid in hdata['rule_ids']: |
| if rid in self.hypergraph.rules: |
| he.rules.append(self.hypergraph.rules[rid]) |
| self.hypergraph.hyperedges[hid] = he |
|
|
| |
| for k, v in state.get('inverted_index', {}).items(): |
| self.hypergraph.inverted_index[k] = set(v) |
|
|
| |
| if hasattr(self, 'mlp_gate') and self.mlp_gate is not None: |
| try: |
| mlp_path = os.path.join(self.persistence_dir, 'mlp_gate.pt') |
| self.mlp_gate.load(mlp_path) |
| except Exception as e: |
| print(f"Failed to load MLP gate: {e}") |
|
|
| print(f"EvoRM state loaded from {state_path}: " |
| f"{len(self.hypergraph.rules)} rules, " |
| f"{len(self.hypergraph.hyperedges)} hyperedges") |
|
|
| except Exception as e: |
| print(f"Failed to load EvoRM state: {e}") |
|
|
|
|
| |
| |
| |
|
|
| if __name__ == "__main__": |
| print("EvoRM Plugin - Self Test") |
| print("=" * 60) |
|
|
| |
| plugin = EvoRMPlugin(client=None) |
|
|
| |
| print("\n1. Testing Rule Encoding...") |
| rationale = """ |
| [DECISIVE] title=same ("locating data sources"), year=same (2003) |
| [SUPPORTING] authors=same |
| """ |
| atoms = plugin.rule_encoding.parse_rationale(rationale, conclusion=1) |
| print(f" Parsed atoms: {[a.to_key() for a in atoms]}") |
|
|
| |
| rule = plugin.rule_encoding.create_rule(rationale, conclusion=1) |
| print(f" Created rule: {rule.rule_id}, conf={rule.conf:.3f}") |
|
|
| |
| print("\n2. Testing Hypergraph Storage...") |
| node_set = {"e_1", "e_2", "a_title", "a_year", "a_authors"} |
| he = plugin.hypergraph.get_or_create_hyperedge( |
| (1, 100), node_set, rule) |
| print(f" Created hyperedge: {he.hyperedge_id}, weight={he.weight:.3f}") |
|
|
| |
| candidates = plugin.hypergraph.get_candidate_rules(["e_1"]) |
| print(f" Candidates for e_1: {len(candidates)} rules") |
|
|
| |
| print("\n3. Testing Stage 1 (Symbolic Filtering)...") |
| es_ctx = {"entity_name": "Test E1", "title": "locating data sources", |
| "year": "2003", "authors": "Smith et al"} |
| et_ctx = {"entity_name": "Test E2", "title": "locating data sources", |
| "year": "2003", "authors": "Smith et al"} |
| decision, triggered, cands = plugin.stage1("e_1", "e_2", es_ctx, et_ctx) |
| print(f" Decision: {decision}, Triggered: {len(triggered)} rules") |
|
|
| |
| print("\n4. Testing Trajectory Recording...") |
| plugin.record_trajectory( |
| "e_1", "e_2", es_ctx, et_ctx, |
| decision=1, rationale=rationale, |
| triggered_rules=triggered, |
| rule_feedback={"rule_1": 0.8}) |
|
|
| |
| print("\n5. Stats:", json.dumps(plugin.get_stats(), indent=2)) |
|
|
| print("\n✅ All tests passed!") |