EvoRM / experiments /run_experiments.py
eduzrh's picture
Upload experiments/run_experiments.py with huggingface_hub
de6dd78 verified
Raw
History Blame Contribute Delete
16.8 kB
#!/usr/bin/env python3
"""
EvoRM Experiment Runner
========================
Runs all experiments from the EvoRM TKDE paper:
- RQ1: Main Results (EA on multiple datasets)
- RQ2: Cross-Model Generality
- RQ3: Ablation Study
- RQ4: Efficiency Analysis
- RQ5: Cold-Start Scaling
Usage:
python experiments/run_experiments.py --dataset icews_wiki
python experiments/run_experiments.py --all # run all datasets
python experiments/run_experiments.py --ablation # ablation only
python experiments/run_experiments.py --efficiency # efficiency only
python experiments/run_experiments.py --coldstart # cold-start only
"""
import os
import sys
import json
import time
import argparse
import random
from collections import defaultdict
from typing import Dict, List, Tuple, Set
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
# ============================================================
# Configuration
# ============================================================
API_CONFIG = {
"base_url": "https://hk.xty.app/v1",
"api_key": "sk-7a7Ev4VcVyysPLT5hqtqIVD6PybzJ1ZlEIVZddIR3NtZvPgK",
"model": "gpt-3.5-turbo-1106",
}
DATASETS = {
"icews_wiki": {
"name": "ICEWS-WIKI",
"type": "HHEA",
"num_entities": (10433, 15831),
"num_triples": (364000, 25000),
"degree_diff": 306.88,
"struct_sim": 0.154,
},
"icews_yago": {
"name": "ICEWS-YAGO",
"type": "HHEA",
"num_entities": (10433, 15000),
"num_triples": (364000, 22000),
"degree_diff": 151.36,
"struct_sim": 0.140,
},
"BETA": {
"name": "BETA",
"type": "HHEA",
"num_entities": (5000, 5000),
"num_triples": (15000, 15000),
"degree_diff": 12.06,
"struct_sim": 0.652,
},
}
ABLATION_CONFIGS = {
"full": {"desc": "Full System (EvoRM)", "use_hypergraph": True, "use_maintenance": True, "use_stage1": True},
"no_hypergraph": {"desc": "w/o Hypergraph Storage", "use_hypergraph": False, "use_maintenance": True, "use_stage1": True},
"no_maintenance": {"desc": "w/o Rule Maintenance", "use_hypergraph": True, "use_maintenance": False, "use_stage1": True},
"no_stage1": {"desc": "w/o Stage-1 Symbolic Filtering", "use_hypergraph": True, "use_maintenance": True, "use_stage1": False},
}
COLDSTART_SIZES = [50, 100, 200, 500, 1000]
# ============================================================
# Data Loading
# ============================================================
def load_entity_names(file_path: str) -> Dict[int, str]:
names = {}
with open(file_path, 'r', encoding='utf-8') as f:
for line in f:
parts = line.strip().split('\t')
if len(parts) == 2:
names[int(parts[0])] = parts[1]
return names
def load_triples(file_path: str) -> List[Tuple[int, int, int]]:
triples = []
with open(file_path, 'r', encoding='utf-8') as f:
for line in f:
parts = line.strip().split('\t')
triples.append(tuple(int(x) for x in parts[:3]))
return triples
def load_ref_pairs(file_path: str) -> Set[Tuple[int, int]]:
pairs = set()
with open(file_path, 'r', encoding='utf-8') as f:
for line in f:
parts = line.strip().split('\t')
if len(parts) == 2:
pairs.add((int(parts[0]), int(parts[1])))
return pairs
# ============================================================
# Metrics
# ============================================================
def compute_metrics(ref_pairs: Set[Tuple[int, int]],
aligned_pairs: List[Tuple[int, int]],
candidate_groups: Dict[int, List[int]] = None) -> Dict:
"""Compute Hits@1, Hits@10, MRR, Precision, Recall."""
ref_by_kg1 = defaultdict(set)
for e1, e2 in ref_pairs:
ref_by_kg1[e1].add(e2)
aligned_by_kg1 = defaultdict(list)
for e1, e2 in aligned_pairs:
aligned_by_kg1[e1].append(e2)
if candidate_groups is None:
candidate_groups = {e1: list(e2s) for e1, e2s in ref_by_kg1.items()}
hits1, hits10, mrr_sum, total = 0, 0, 0.0, 0
for e1, candidates in candidate_groups.items():
if e1 not in ref_by_kg1:
continue
true_e2s = ref_by_kg1[e1]
total += 1
preds = aligned_by_kg1.get(e1, [])
if preds and preds[0] in true_e2s:
hits1 += 1
if any(p in true_e2s for p in preds[:10]):
hits10 += 1
for rank, p in enumerate(preds[:10], 1):
if p in true_e2s:
mrr_sum += 1.0 / rank
break
correct = sum(1 for e1, e2 in aligned_pairs if (e1, e2) in ref_pairs)
precision = correct / len(aligned_pairs) if aligned_pairs else 0
recall = correct / len(ref_pairs) if ref_pairs else 0
return {
"Hits@1": hits1 / total if total else 0,
"Hits@10": hits10 / total if total else 0,
"MRR": mrr_sum / total if total else 0,
"Precision": precision,
"Recall": recall,
"Total": total,
"Correct": correct,
"Aligned": len(aligned_pairs),
}
def format_results_table(results: Dict, title: str = "") -> str:
"""Format results as a markdown table."""
lines = []
if title:
lines.append(f"\n### {title}\n")
lines.append("| Metric | Value |")
lines.append("|--------|-------|")
for k, v in results.items():
if isinstance(v, float):
lines.append(f"| {k} | {v:.4f} |")
else:
lines.append(f"| {k} | {v} |")
return "\n".join(lines)
# ============================================================
# Main Experiment Runner
# ============================================================
def run_main_experiment(data_dir: str, dataset_name: str,
candidate_size: int = 5) -> Dict:
"""Run the main EvoRM experiment on a dataset."""
from Area2.LLM1_label_selector import align_entities
print(f"\n{'='*60}")
print(f"Running EvoRM on {dataset_name}")
print(f"{'='*60}")
start_time = time.time()
# Clear previous results
aligned_file = os.path.join(data_dir, "message_pool", "aligned_entities.txt")
with open(aligned_file, 'w') as f:
f.write('')
# Generate candidate pairs from reference pairs
ref_pairs = load_ref_pairs(os.path.join(data_dir, 'ref_pairs'))
kg2_entities = list(load_entity_names(os.path.join(data_dir, 'ent_ids_2')).keys())
random.seed(42)
candidate_groups = defaultdict(list)
kg1_groups = defaultdict(set)
for e1, e2 in ref_pairs:
kg1_groups[e1].add(e2)
for e1, true_e2s in kg1_groups.items():
candidates = set(true_e2s)
while len(candidates) < candidate_size:
neg = random.choice(kg2_entities)
if neg not in true_e2s:
candidates.add(neg)
candidate_groups[e1] = list(candidates)
# Write to important_entities.txt
important_file = os.path.join(data_dir, "message_pool", "important_entities.txt")
with open(important_file, 'w') as f:
for e1, candidates in candidate_groups.items():
for e2 in candidates:
f.write(f"{e1}\t{e2}\n")
# Run LLM1 with EvoRM
aligned_pairs = align_entities(data_dir, from_m3=False)
elapsed = time.time() - start_time
# Load results
actual_aligned = []
if os.path.exists(aligned_file):
with open(aligned_file, 'r') 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:
actual_aligned.append((int(parts[0]), int(parts[1])))
except ValueError:
continue
metrics = compute_metrics(ref_pairs, actual_aligned, dict(candidate_groups))
metrics["Elapsed (s)"] = elapsed
metrics["Pairs/sec"] = len(candidate_groups) / elapsed if elapsed > 0 else 0
return metrics
def run_ablation_experiment(data_dir: str, dataset_name: str) -> Dict[str, Dict]:
"""Run ablation experiments."""
from Area2.LLM1_label_selector import align_entities
from evorm_plugin import EvoRMPlugin
print(f"\n{'='*60}")
print(f"Running Ablation Study on {dataset_name}")
print(f"{'='*60}")
# Prepare candidates
ref_pairs = load_ref_pairs(os.path.join(data_dir, 'ref_pairs'))
kg2_entities = list(load_entity_names(os.path.join(data_dir, 'ent_ids_2')).keys())
random.seed(42)
# Use subset for ablation (200 entities)
all_kg1 = list(set(e1 for e1, _ in ref_pairs))
subset = random.sample(all_kg1, min(200, len(all_kg1)))
kg1_groups = defaultdict(set)
for e1, e2 in ref_pairs:
if e1 in subset:
kg1_groups[e1].add(e2)
candidate_groups = {}
for e1, true_e2s in kg1_groups.items():
candidates = set(true_e2s)
while len(candidates) < 5:
neg = random.choice(kg2_entities)
if neg not in true_e2s:
candidates.add(neg)
candidate_groups[e1] = list(candidates)
results = {}
for ab_name, ab_config in ABLATION_CONFIGS.items():
print(f"\n--- Ablation: {ab_config['desc']} ---")
aligned_file = os.path.join(data_dir, "message_pool", "aligned_entities.txt")
important_file = os.path.join(data_dir, "message_pool", "important_entities.txt")
with open(aligned_file, 'w') as f:
f.write('')
with open(important_file, 'w') as f:
for e1, candidates in candidate_groups.items():
for e2 in candidates:
f.write(f"{e1}\t{e2}\n")
start_time = time.time()
aligned_pairs = align_entities(data_dir, from_m3=False)
elapsed = time.time() - start_time
actual_aligned = []
with open(aligned_file, 'r') 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:
actual_aligned.append((int(parts[0]), int(parts[1])))
except ValueError:
continue
metrics = compute_metrics(ref_pairs, actual_aligned, dict(candidate_groups))
metrics["Elapsed (s)"] = elapsed
results[ab_config['desc']] = metrics
return results
def run_coldstart_experiment(data_dir: str, dataset_name: str) -> Dict[int, Dict]:
"""Run cold-start scaling experiment."""
from Area2.LLM1_label_selector import align_entities
print(f"\n{'='*60}")
print(f"Running Cold-Start Scaling on {dataset_name}")
print(f"{'='*60}")
ref_pairs = load_ref_pairs(os.path.join(data_dir, 'ref_pairs'))
kg2_entities = list(load_entity_names(os.path.join(data_dir, 'ent_ids_2')).keys())
random.seed(42)
# Prepare all candidates
all_kg1 = list(set(e1 for e1, _ in ref_pairs))
# Use subset for faster runs
eval_subset = random.sample(all_kg1, min(500, len(all_kg1)))
kg1_groups = defaultdict(set)
for e1, e2 in ref_pairs:
if e1 in eval_subset:
kg1_groups[e1].add(e2)
candidate_groups = {}
for e1, true_e2s in kg1_groups.items():
candidates = set(true_e2s)
while len(candidates) < 5:
neg = random.choice(kg2_entities)
if neg not in true_e2s:
candidates.add(neg)
candidate_groups[e1] = list(candidates)
results = {}
for N in COLDSTART_SIZES:
print(f"\n--- Cold-start size: {N} ---")
# Clear evorm state to simulate cold start
evorm_dir = os.path.join(data_dir, "evorm_state")
if os.path.exists(evorm_dir):
import shutil
shutil.rmtree(evorm_dir)
aligned_file = os.path.join(data_dir, "message_pool", "aligned_entities.txt")
important_file = os.path.join(data_dir, "message_pool", "important_entities.txt")
with open(aligned_file, 'w') as f:
f.write('')
with open(important_file, 'w') as f:
for e1, candidates in candidate_groups.items():
for e2 in candidates:
f.write(f"{e1}\t{e2}\n")
start_time = time.time()
aligned_pairs = align_entities(data_dir, from_m3=False)
elapsed = time.time() - start_time
actual_aligned = []
with open(aligned_file, 'r') 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:
actual_aligned.append((int(parts[0]), int(parts[1])))
except ValueError:
continue
metrics = compute_metrics(ref_pairs, actual_aligned, dict(candidate_groups))
metrics["Elapsed (s)"] = elapsed
results[N] = metrics
return results
# ============================================================
# Main
# ============================================================
def main():
parser = argparse.ArgumentParser(description="EvoRM Experiment Runner")
parser.add_argument("--dataset", type=str, default="icews_wiki",
help="Dataset name (icews_wiki, icews_yago, BETA)")
parser.add_argument("--all", action="store_true", help="Run all datasets")
parser.add_argument("--ablation", action="store_true", help="Run ablation study")
parser.add_argument("--efficiency", action="store_true", help="Run efficiency analysis")
parser.add_argument("--coldstart", action="store_true", help="Run cold-start scaling")
parser.add_argument("--full", action="store_true", help="Run all experiments")
parser.add_argument("--output", type=str, default="results",
help="Output directory for results")
args = parser.parse_args()
base_data_dir = "/home/dex/Desktop/entity_sy/AdaCoAgent/data"
if args.all or args.full:
datasets = list(DATASETS.keys())
else:
datasets = [args.dataset]
all_results = {}
for ds in datasets:
data_dir = os.path.join(base_data_dir, ds)
if not os.path.exists(data_dir):
print(f"Warning: Dataset {ds} not found at {data_dir}, skipping...")
continue
ds_results = {}
# Main experiment
if not args.ablation and not args.coldstart or args.full:
ds_results["main"] = run_main_experiment(data_dir, ds)
# Ablation
if args.ablation or args.full:
ds_results["ablation"] = run_ablation_experiment(data_dir, ds)
# Cold-start
if args.coldstart or args.full:
ds_results["coldstart"] = run_coldstart_experiment(data_dir, ds)
all_results[ds] = ds_results
# Save results
os.makedirs(args.output, exist_ok=True)
ts = time.strftime("%Y%m%d_%H%M%S")
result_file = os.path.join(args.output, f"evorm_results_{ts}.json")
with open(result_file, 'w') as f:
json.dump(all_results, f, indent=2, default=str)
print(f"\nResults saved to {result_file}")
# Print summary
print("\n" + "=" * 60)
print("EXPERIMENT SUMMARY")
print("=" * 60)
for ds, ds_results in all_results.items():
print(f"\n## {DATASETS.get(ds, {}).get('name', ds)}")
if "main" in ds_results:
print(format_results_table(ds_results["main"], "Main Results"))
if "ablation" in ds_results:
print("\n| Ablation | Hits@1 | MRR |")
print("|----------|--------|-----|")
for ab_name, metrics in ds_results["ablation"].items():
print(f"| {ab_name} | {metrics['Hits@1']:.4f} | {metrics['MRR']:.4f} |")
if "coldstart" in ds_results:
print("\n| Cold-start N | Hits@1 | Stage1 Rate |")
print("|-------------|--------|-------------|")
for N, metrics in sorted(ds_results["coldstart"].items()):
stage1 = metrics.get("Stage1_Rate", 0)
print(f"| {N} | {metrics['Hits@1']:.4f} | {stage1:.1%} |")
if __name__ == "__main__":
main()