eduzrh's picture
Upload code/baselines/evorm_wrappers/zero_cot.py with huggingface_hub
633a5b1 verified
Raw
History Blame Contribute Delete
10.5 kB
#!/usr/bin/env python3
"""
ZeroCoT Baseline: Zero-shot Chain-of-Thought Entity Alignment
=============================================================
Paper Table II: ZeroCoT* achieves Hits@1=0.639 on DBP15K without EvoRM,
and 0.711 (+11.3%) with EvoRM.
This baseline uses a simple zero-shot CoT prompt: "Let's think step by step"
followed by the entity matching decision. No reasoning framework, no rules.
Supports:
- Standalone mode (baseline)
- +EvoRM mode (plug-and-play enhancement)
"""
import os
import queue
import threading
import sys
from tqdm import tqdm
from openai import OpenAI
from collections import defaultdict
import json
sys.path.append('/root/autodl-tmp/AdaCoAgentEA')
from ThreadPoolExecutor import ThreadPoolExecutor
import tokens_cal
def load_entity_names(file_path):
entity_names = {}
with open(file_path, 'r', encoding='utf-8') as f:
for line in f:
parts = line.strip().split('\t')
if len(parts) == 2:
entity_names[int(parts[0])] = parts[1]
return entity_names
def load_triples(file_path):
triples = []
with open(file_path, 'r', encoding='utf-8') as f:
for line in f:
parts = line.strip().split('\t')
triples.append([int(x) for x in parts[:3]])
return triples
def get_entity_context(entity_id, entity_names, triples, rel_names, n=3):
relations = []
for h, r, t in triples:
if h == entity_id:
rel_str = rel_names.get(r, str(r))
tail_str = entity_names.get(t, str(t))
relations.append(f"- Has relation '{rel_str}' with {tail_str}")
elif t == entity_id:
rel_str = rel_names.get(r, str(r))
head_str = entity_names.get(h, str(h))
relations.append(f"- Is {rel_str} of {head_str}")
if len(relations) >= n:
break
context = f"Entity Name: {entity_names.get(entity_id, 'Unknown')}\n"
context += "Relationships:\n" + "\n".join(relations[:n])
return context
def group_candidates(input_file):
groups = defaultdict(list)
with open(input_file, 'r', encoding='utf-8') as f:
for line in f:
parts = line.strip().split('\t')
if len(parts) == 2:
try:
e1, e2 = int(parts[0]), int(parts[1])
groups[e1].append(e2)
except ValueError:
continue
return groups
def deduplicate_output_file(file_path):
if not os.path.exists(file_path):
return
unique_pairs = set()
with open(file_path, 'r', encoding='utf-8') as f:
for line in f:
line = line.strip()
if not line:
continue
parts = line.split('\t')
if len(parts) == 2 and parts[0] and parts[1]:
try:
e1, e2 = int(parts[0]), int(parts[1])
unique_pairs.add((e1, e2))
except ValueError:
continue
with open(file_path, 'w', encoding='utf-8') as f:
for e1, e2 in sorted(unique_pairs):
f.write(f"{e1}\t{e2}\n")
print(f"Deduplicated: {len(unique_pairs)} unique pairs")
# Zero-shot CoT prompt
ZERO_COT_PROMPT = """You are a knowledge graph entity alignment expert. Determine if two entities represent the same real-world object.
Think step by step:
1. Compare the entity names.
2. Compare the relationships of each entity.
3. Consider if the relationships describe the same real-world connections.
4. Make a decision.
Entity A (ID: {kg1_id}):
{context_a}
Entity B (ID: {kg2_id}):
{context_b}
Let's think step by step. After your analysis, output ONLY:
- The entity ID of the matching candidate (e.g., "{kg2_id}") if they match
- "No" if they do not match
Your response:"""
def align_entities(data_dir, from_m3=False, ablation_config=None,
no_optimization_tool=False, use_evorm=False,
evorm_ablation_mode=None):
"""
ZeroCoT Entity Alignment.
Args:
data_dir: Path to dataset directory
use_evorm: If True, wrap with EvoRM plugin
evorm_ablation_mode: EvoRM ablation mode (if use_evorm=True)
Returns:
List of aligned pairs [(kg1_id, kg2_id), ...]
"""
client = OpenAI(
base_url="https://hk.xty.app/v1",
api_key="sk-7a7Ev4VcVyysPLT5hqtqIVD6PybzJ1ZlEIVZddIR3NtZvPgK",
)
# EvoRM plugin
evorm_plugin = None
if use_evorm:
from evorm_plugin import EvoRMPlugin
evorm_plugin = EvoRMPlugin(
client=client,
ablation_mode=evorm_ablation_mode,
persistence_dir=os.path.join(data_dir, 'evorm_state_zero_cot'),
)
print(f"ZeroCoT + EvoRM enabled (ablation={evorm_ablation_mode})")
ent_names_1 = load_entity_names(os.path.join(data_dir, 'ent_ids_1'))
ent_names_2 = load_entity_names(os.path.join(data_dir, 'ent_ids_2'))
rel_names_1 = load_entity_names(os.path.join(data_dir, 'rel_ids_1'))
rel_names_2 = load_entity_names(os.path.join(data_dir, 'rel_ids_2'))
triples_1 = load_triples(os.path.join(data_dir, 'triples_1'))
triples_2 = load_triples(os.path.join(data_dir, 'triples_2'))
LLM1_PRIVATE_MESSAGE_POOL = {
'important_entities': os.path.join(data_dir, "message_pool", "important_entities.txt"),
'ucon_similarity_results': os.path.join(data_dir, "message_pool", "ucon_similarity_results.txt"),
'aligned_entities': os.path.join(data_dir, "message_pool", "aligned_entities.txt"),
}
input_file = LLM1_PRIVATE_MESSAGE_POOL['ucon_similarity_results'] if from_m3 else LLM1_PRIVATE_MESSAGE_POOL['important_entities']
output_file = LLM1_PRIVATE_MESSAGE_POOL['aligned_entities']
candidate_groups = group_candidates(input_file)
aligned_pairs = []
lock = threading.Lock()
executor = ThreadPoolExecutor(max_workers=5)
result_queue = queue.Queue()
def zero_cot_task(kg1_entity, kg2_candidates):
try:
context_a = get_entity_context(kg1_entity, ent_names_1, triples_1, rel_names_1)
for kg2_entity in kg2_candidates:
context_b = get_entity_context(kg2_entity, ent_names_2, triples_2, rel_names_2)
# Build prompt
prompt = ZERO_COT_PROMPT.format(
kg1_id=kg1_entity, kg2_id=kg2_entity,
context_a=context_a, context_b=context_b,
)
if evorm_plugin is not None:
# Build context dicts
es_ctx = {'entity_name': ent_names_1.get(kg1_entity, '')}
et_ctx = {'entity_name': ent_names_2.get(kg2_entity, '')}
for h, r, t in triples_1:
if h == kg1_entity:
rel = rel_names_1.get(r, str(r))
neighbor = ent_names_1.get(t, str(t))
key = f"neighbors_{rel}"
es_ctx.setdefault(key, set()).add(neighbor)
for h, r, t in triples_2:
if t == kg2_entity:
rel = rel_names_2.get(r, str(r))
neighbor = ent_names_2.get(h, str(h))
key = f"neighbors_{rel}"
et_ctx.setdefault(key, set()).add(neighbor)
# Stage 1: symbolic filtering
decision, triggered, candidates = evorm_plugin.stage1(
str(kg1_entity), str(kg2_entity), es_ctx, et_ctx)
if decision is not None:
if decision == 1:
with lock:
result_queue.put((kg1_entity, kg2_entity))
aligned_pairs.append((kg1_entity, kg2_entity))
continue
# Stage 2: LLM with evidence
result = evorm_plugin.stage2(
str(kg1_entity), str(kg2_entity), es_ctx, et_ctx,
triggered, candidates, prompt)
decision = result.get('decision')
if decision == 1:
with lock:
result_queue.put((kg1_entity, kg2_entity))
aligned_pairs.append((kg1_entity, kg2_entity))
evorm_plugin.record_trajectory(
str(kg1_entity), str(kg2_entity), es_ctx, et_ctx,
decision=decision or 0,
rationale=result.get('rationale', ''),
triggered_rules=triggered,
rule_feedback=result.get('rule_feedback', {}))
else:
# Baseline: direct LLM call
response = client.chat.completions.create(
model="gpt-3.5-turbo-1106",
messages=[{'role': 'user', 'content': prompt}],
temperature=0.1,
)
answer = response.choices[0].message.content.strip()
tokens_cal.update_add_var(response.usage.total_tokens)
if answer.lower() != "no" and str(kg2_entity) in answer:
with lock:
result_queue.put((kg1_entity, kg2_entity))
aligned_pairs.append((kg1_entity, kg2_entity))
break
except Exception as e:
print(f"Error processing entity {kg1_entity}: {str(e)}")
for kg1_entity_c, kg2_candidates_c in tqdm(candidate_groups.items(), desc="ZeroCoT"):
executor.submit(zero_cot_task, kg1_entity_c, kg2_candidates_c)
executor.shutdown(wait=True)
with open(output_file, 'a+', encoding='utf-8') as output_f:
while not result_queue.empty():
kg1_entity, kg2_id = result_queue.get()
output_f.write(f"{kg1_entity}\t{kg2_id}\n")
output_f.flush()
deduplicate_output_file(output_file)
if evorm_plugin:
stats = evorm_plugin.get_stats()
print(f"ZeroCoT+EvoRM stats: {json.dumps(stats, indent=2, default=str)}")
return aligned_pairs
if __name__ == "__main__":
data_dir = "/root/autodl-tmp/AdaCoAgentEA/data/icews_wiki"
print("ZeroCoT Baseline (no EvoRM):")
pairs = align_entities(data_dir, use_evorm=False)
print(f"Found {len(pairs)} aligned pairs.")