-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathphase1_data_collection.py
More file actions
187 lines (156 loc) · 6.65 KB
/
Copy pathphase1_data_collection.py
File metadata and controls
187 lines (156 loc) · 6.65 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
from SPARQLWrapper import SPARQLWrapper, JSON
import json
import time
import os
import sqlite3
import pandas as pd
from datasets import load_dataset
from sklearn.model_selection import train_test_split
from dotenv import load_dotenv
load_dotenv()
DB_NAME = "temphal.db"
def get_wikidata_sparql():
sparql = SPARQLWrapper('https://query.wikidata.org/sparql')
sparql.addCustomHttpHeader('User-Agent', 'TempHal-Research/1.0 (student project)')
sparql.setReturnFormat(JSON)
return sparql
def extract_volatile_claims(limit_per_query=1000, total_limit=3000):
"""Extracts volatile claims (CEOs) from WikiData."""
sparql = get_wikidata_sparql()
QUERY = '''
SELECT ?item ?itemLabel ?ceo ?ceoLabel ?start ?end WHERE {{
?item p:P169 ?stmt .
?stmt ps:P169 ?ceo .
OPTIONAL {{ ?stmt pq:P580 ?start . }}
OPTIONAL {{ ?stmt pq:P582 ?end . }}
SERVICE wikibase:label {{
bd:serviceParam wikibase:language 'en'. }}
}} LIMIT {limit} OFFSET {offset}
'''
volatile_claims = []
for offset in range(0, total_limit, limit_per_query):
print(f"Fetching volatile claims (offset {offset})...")
sparql.setQuery(QUERY.format(limit=limit_per_query, offset=offset))
try:
res = sparql.query().convert()
for r in res['results']['bindings']:
item_label = r['itemLabel']['value']
ceo_label = r['ceoLabel']['value']
# Skip if labels are same or generic
if item_label == ceo_label or "Q" in item_label:
continue
row = {
'claim': f"{ceo_label} is the CEO of {item_label}",
'temporal_type': 'volatile',
'entity_type': 'organization',
'ground_truth_date': r.get('start', {}).get('value', ''),
'source': 'wikidata_p169',
'volatility_score': 0.8
}
volatile_claims.append(row)
except Exception as e:
print(f'Offset {offset} failed: {e}')
time.sleep(2)
print(f"Extracted {len(volatile_claims)} volatile claims.")
return volatile_claims
def extract_stable_claims(limit_per_query=500, total_limit=1000):
"""Extracts stable claims (Birth dates of historical figures) from WikiData."""
sparql = get_wikidata_sparql()
# Birth dates of people who died before 2020
QUERY = '''
SELECT ?item ?itemLabel ?birth ?death WHERE {{
?item wdt:P31 wd:Q5 . # human
?item wdt:P569 ?birth .
?item wdt:P570 ?death .
FILTER(?death < "2020-01-01T00:00:00Z"^^xsd:dateTime)
SERVICE wikibase:label {{
bd:serviceParam wikibase:language 'en'. }}
}} LIMIT {limit} OFFSET {offset}
'''
stable_claims = []
for offset in range(0, total_limit, limit_per_query):
print(f"Fetching stable claims (offset {offset})...")
sparql.setQuery(QUERY.format(limit=limit_per_query, offset=offset))
try:
res = sparql.query().convert()
for r in res['results']['bindings']:
item_label = r['itemLabel']['value']
date_val = r['birth']['value']
# Format claim
year = date_val.split('-')[0].replace('+', '')
row = {
'claim': f"{item_label} was born in {year}",
'temporal_type': 'stable',
'entity_type': 'person',
'ground_truth_date': date_val,
'source': 'wikidata_p569',
'volatility_score': 0.1
}
stable_claims.append(row)
except Exception as e:
print(f'Offset {offset} failed: {e}')
time.sleep(2)
print(f"Extracted {len(stable_claims)} stable claims.")
return stable_claims
def load_external_data():
"""Loads claims from TruthfulQA."""
print("Loading TruthfulQA dataset...")
try:
tqa = load_dataset('truthful_qa', 'generation', split='validation')
uncertain_keywords = ['current', 'now', 'today', 'presently', 'latest', 'moment', 'who is', 'what is']
tqa_claims = [
{'claim': r['best_answer'],
'temporal_type': 'uncertain',
'source': 'truthfulqa',
'volatility_score': 0.5}
for r in tqa
if any(kw in r['question'].lower() for kw in uncertain_keywords)
]
except Exception as e:
print(f"Error loading TruthfulQA: {e}")
tqa_claims = []
# Skipping FEVER due to compatibility issues with newer 'datasets' versions
fever_claims = []
print(f"Loaded {len(fever_claims)} FEVER claims (skipped) and {len(tqa_claims)} TruthfulQA claims.")
return fever_claims, tqa_claims
def finalize_and_store(all_claims):
"""Splits and stores the claims in the SQLite database."""
df = pd.DataFrame(all_claims)
# Ensure we have enough data for the targets
# Target distribution: 40% stable, 40% volatile, 20% uncertain
# Total ~2000
volatile_df = df[df['temporal_type'] == 'volatile'].sample(min(800, len(df[df['temporal_type'] == 'volatile'])))
stable_df = df[df['temporal_type'] == 'stable'].sample(min(800, len(df[df['temporal_type'] == 'stable'])))
uncertain_df = df[df['temporal_type'] == 'uncertain'].sample(min(400, len(df[df['temporal_type'] == 'uncertain'])))
final_df = pd.concat([volatile_df, stable_df, uncertain_df])
print(f"Final dataset size: {len(final_df)}")
print(final_df['temporal_type'].value_counts())
# Stratified split
train_df, test_val_df = train_test_split(
final_df, test_size=0.4, stratify=final_df['temporal_type'], random_state=42
)
val_df, test_df = train_test_split(
test_val_df, test_size=0.5, stratify=test_val_df['temporal_type'], random_state=42
)
train_df['split'] = 'train'
val_df['split'] = 'validation'
test_df['split'] = 'test'
final_with_splits = pd.concat([train_df, val_df, test_df])
# Store in SQLite
conn = sqlite3.connect(DB_NAME)
final_with_splits.to_sql('claims', conn, if_exists='append', index=False)
conn.close()
print("Claims stored in SQLite successfully.")
def main():
# 1. WikiData Volatile
volatile = extract_volatile_claims(total_limit=3000)
# 2. WikiData Stable
stable_wd = extract_stable_claims(total_limit=2000)
# 3. External Data
fever, tqa = load_external_data()
# Combine
all_claims = volatile + stable_wd + fever + tqa
# 4. Finalize and store
finalize_and_store(all_claims)
if __name__ == "__main__":
main()