| |
| """ |
| Preprocess new datasets for EvoRM experiments: |
| - DBP-WIKI (EA): D-W-15K and D-W-100K |
| - ZESHEL (EL): 4 test domains |
| - MIMIC (SM): OMAP clinical schema matching |
| - Synthea (SM): OMAP clinical schema matching |
| - WDC (ER): product matching |
| """ |
|
|
| import os |
| import sys |
| import json |
| import random |
| from collections import defaultdict |
|
|
| DATA_DIR = '/root/autodl-tmp/AdaCoAgentEA/data' |
| random.seed(42) |
|
|
| |
| |
| |
| def prepare_dbp_wiki(): |
| """Convert OpenEA D-W datasets to EA format for AdaCoAgentEA.""" |
| src_dir = os.path.join(DATA_DIR, 'ea', 'dbp_wiki', 'OpenEA_dataset_v2.0') |
| out_dir = os.path.join(DATA_DIR, 'ea', 'dbp_wiki_processed') |
| os.makedirs(out_dir, exist_ok=True) |
| |
| for variant in ['D_W_15K_V1', 'D_W_15K_V2']: |
| vdir = os.path.join(src_dir, variant) |
| if not os.path.exists(vdir): |
| continue |
| |
| |
| ent_links = [] |
| with open(os.path.join(vdir, 'ent_links'), 'r') as f: |
| for line in f: |
| parts = line.strip().split('\t') |
| if len(parts) == 2: |
| ent_links.append((parts[0], parts[1])) |
| |
| print(f" {variant}: {len(ent_links)} entity links") |
| |
| |
| for fold in range(1, 6): |
| fold_dir = os.path.join(vdir, '721_5fold', str(fold)) |
| test_set = set() |
| valid_set = set() |
| |
| for fname, linkset in [('test_links', test_set), ('valid_links', valid_set)]: |
| fpath = os.path.join(fold_dir, fname) |
| if os.path.exists(fpath): |
| with open(fpath, 'r') as f: |
| for line in f: |
| parts = line.strip().split('\t') |
| if len(parts) == 2: |
| linkset.add((parts[0], parts[1])) |
| |
| |
| test_pairs = [] |
| for e1, e2 in ent_links: |
| if (e1, e2) in test_set: |
| test_pairs.append({ |
| 'entity_a': e1.split('/')[-1], |
| 'entity_b': e2.split('/')[-1], |
| 'label': 1, |
| }) |
| elif (e1, e2) in valid_set: |
| |
| test_pairs.append({ |
| 'entity_a': e1.split('/')[-1], |
| 'entity_b': e2.split('/')[-1], |
| 'label': 1, |
| }) |
| |
| |
| |
| neg_count = 0 |
| max_neg = min(len(test_pairs) * 2, 500) |
| all_ent_links_set = set(ent_links) |
| for e1, e2 in ent_links: |
| if neg_count >= max_neg: |
| break |
| if (e1, e2) not in test_set and (e1, e2) not in valid_set: |
| test_pairs.append({ |
| 'entity_a': e1.split('/')[-1], |
| 'entity_b': e2.split('/')[-1], |
| 'label': 0, |
| }) |
| neg_count += 1 |
| |
| random.shuffle(test_pairs) |
| |
| out_path = os.path.join(out_dir, f'{variant}_fold{fold}.json') |
| with open(out_path, 'w') as f: |
| json.dump(test_pairs, f) |
| pos = sum(1 for p in test_pairs if p['label'] == 1) |
| print(f" fold{fold}: {len(test_pairs)} pairs ({pos} pos)") |
| |
| print("DBP-WIKI: Done") |
|
|
|
|
| |
| |
| |
| def prepare_zeshel(): |
| """Convert ZESHEL to EL format for LELA baseline.""" |
| data_dir = os.path.join(DATA_DIR, 'el', 'zeshel') |
| out_dir = os.path.join(DATA_DIR, 'el', 'zeshel_processed') |
| os.makedirs(out_dir, exist_ok=True) |
| |
| |
| print(" Loading ZESHEL dictionary...") |
| entities = defaultdict(dict) |
| with open(os.path.join(data_dir, 'dictionary', 'dictionary-00001-of-00001.jsonl')) as f: |
| for line in f: |
| d = json.loads(line) |
| entities[d['subset']][d['id']] = { |
| 'name': d['name'], |
| 'description': d['description'][:300] if d['description'] else '', |
| } |
| print(f" Loaded {sum(len(v) for v in entities.values())} entities across {len(entities)} domains") |
| |
| |
| test_data = [] |
| with open(os.path.join(data_dir, 'data', 'test-00001-of-00001.jsonl')) as f: |
| for line in f: |
| test_data.append(json.loads(line)) |
| |
| |
| by_domain = defaultdict(list) |
| for item in test_data: |
| by_domain[item['subset']].append(item) |
| |
| test_domains = ['forgotten_realms', 'lego', 'star_trek', 'yugioh'] |
| |
| for domain in test_domains: |
| domain_items = by_domain.get(domain, []) |
| if not domain_items: |
| print(f" WARNING: {domain} has no test items") |
| continue |
| |
| el_items = [] |
| for item in domain_items[:500]: |
| |
| candidates = [] |
| for ent_info in item.get('entities', []): |
| ent_id = ent_info.get('id', '') |
| if ent_id and ent_id in entities[domain]: |
| candidates.append(entities[domain][ent_id]) |
| |
| if not candidates: |
| continue |
| |
| |
| gold_is_first = len(item.get('entities', [])) > 0 |
| if len(candidates) > 10: |
| candidates = candidates[:10] |
| |
| el_items.append({ |
| 'mention': item.get('text', '')[:200], |
| 'context': item.get('text', ''), |
| 'candidates': candidates, |
| 'gold': 'C1', |
| 'domain': domain, |
| }) |
| |
| out_path = os.path.join(out_dir, f'zeshel_{domain}.json') |
| with open(out_path, 'w') as f: |
| json.dump(el_items, f) |
| print(f" ZESHEL {domain}: {len(el_items)} EL items") |
| |
| print("ZESHEL: Done") |
|
|
|
|
| |
| |
| |
| def prepare_smat(): |
| """Convert SMAT OMAP data to SM format.""" |
| import openpyxl |
| |
| data_dir = os.path.join(DATA_DIR, 'sm', 'synthea', 'datasets', 'omap') |
| out_dir = os.path.join(DATA_DIR, 'sm', 'smat_processed') |
| os.makedirs(out_dir, exist_ok=True) |
| |
| for name, fname in [('synthea', 'omop_synthea_data.xlsx'), |
| ('mimic', 'omop_mimic_data.xlsx')]: |
| wb = openpyxl.load_workbook(os.path.join(data_dir, fname)) |
| ws = wb.active |
| |
| items = [] |
| for row in ws.iter_rows(min_row=2, values_only=True): |
| omop, table, des1, des2, label, d1, d2, d3, d4 = row |
| label = int(label) if label is not None else 0 |
| items.append({ |
| 'source': { |
| 'table': str(omop or '').split('-')[0] if omop else '', |
| 'column': str(omop or '').split('-')[1] if omop and '-' in str(omop) else str(omop or ''), |
| 'name': str(omop or ''), |
| 'description': str(des1 or ''), |
| }, |
| 'target': { |
| 'table': str(table or '').split('-')[0] if table else '', |
| 'column': str(table or '').split('-')[1] if table and '-' in str(table) else str(table or ''), |
| 'name': str(table or ''), |
| 'description': str(des2 or ''), |
| }, |
| 'label': label, |
| }) |
| |
| pos_items = [it for it in items if it['label'] == 1] |
| neg_items = [it for it in items if it['label'] == 0] |
| |
| |
| max_neg = min(len(neg_items), len(pos_items) * 5) |
| if max_neg < len(neg_items): |
| neg_items = random.sample(neg_items, max_neg) |
| |
| balanced = pos_items + neg_items |
| random.shuffle(balanced) |
| |
| out_path = os.path.join(out_dir, f'{name}.json') |
| with open(out_path, 'w') as f: |
| json.dump(balanced, f) |
| print(f" SMAT {name}: {len(pos_items)} pos, {len(neg_items)} neg → {len(balanced)} total") |
| |
| print("SMAT: Done") |
|
|
|
|
| |
| |
| |
| def prepare_wdc(): |
| """Convert WDC multi-class to pair-wise ER format.""" |
| import pandas as pd |
| |
| data_dir = os.path.join(DATA_DIR, 'er', 'wdc') |
| out_dir = os.path.join(DATA_DIR, 'er', 'wdc_processed') |
| os.makedirs(out_dir, exist_ok=True) |
| |
| parquet_path = os.path.join(data_dir, 'data', 'train-00000-of-00001.parquet') |
| if not os.path.exists(parquet_path): |
| print("WDC: Parquet file not found, skipping") |
| return |
| |
| df = pd.read_parquet(parquet_path) |
| print(f" WDC: {len(df)} products, {df['cluster_id'].nunique()} clusters") |
| |
| |
| clusters = defaultdict(list) |
| for _, row in df.iterrows(): |
| clusters[row['cluster_id']].append(row.to_dict()) |
| |
| |
| pairs = [] |
| cluster_ids = list(clusters.keys()) |
| |
| |
| for cid, items in clusters.items(): |
| if len(items) >= 2: |
| for i in range(min(len(items), 5)): |
| for j in range(i+1, min(len(items), 5)): |
| pairs.append({ |
| 'record_a': {k: str(v) for k, v in items[i].items() if k != 'cluster_id'}, |
| 'record_b': {k: str(v) for k, v in items[j].items() if k != 'cluster_id'}, |
| 'label': 1, |
| }) |
| |
| |
| neg_count = 0 |
| max_neg = min(len(pairs), 500) |
| while neg_count < max_neg: |
| c1, c2 = random.sample(cluster_ids, 2) |
| if clusters[c1] and clusters[c2]: |
| p1 = random.choice(clusters[c1]) |
| p2 = random.choice(clusters[c2]) |
| pairs.append({ |
| 'record_a': {k: str(v) for k, v in p1.items() if k != 'cluster_id'}, |
| 'record_b': {k: str(v) for k, v in p2.items() if k != 'cluster_id'}, |
| 'label': 0, |
| }) |
| neg_count += 1 |
| |
| random.shuffle(pairs) |
| |
| out_path = os.path.join(out_dir, 'wdc_er.json') |
| with open(out_path, 'w') as f: |
| json.dump(pairs, f) |
| pos = sum(1 for p in pairs if p['label']==1) |
| print(f" WDC: {len(pairs)} pairs ({pos} pos, {len(pairs)-pos} neg)") |
| |
| print("WDC: Done") |
|
|
|
|
| if __name__ == '__main__': |
| print("=" * 60) |
| print("Preparing New Datasets for EvoRM Experiments") |
| print("=" * 60) |
| |
| print("\n1. DBP-WIKI (Entity Alignment)") |
| prepare_dbp_wiki() |
| |
| print("\n2. ZESHEL (Entity Linking)") |
| prepare_zeshel() |
| |
| print("\n3. SMAT (Schema Matching)") |
| prepare_smat() |
| |
| print("\n4. WDC (Entity Resolution)") |
| prepare_wdc() |
| |
| print("\n" + "=" * 60) |
| print("All datasets prepared!") |
|
|