relrae 0.9.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- relrae/RELRaE.py +91 -0
- relrae/__init__.py +6 -0
- relrae/__main__.py +4 -0
- relrae/_version.py +3 -0
- relrae/cli.py +65 -0
- relrae/modules/LLM_Ref.py +525 -0
- relrae/modules/RuBREx.py +315 -0
- relrae/modules/__init__.py +1 -0
- relrae/modules/human_fix.py +248 -0
- relrae/relrae_components/config/LLM_ref_conf.txt +14 -0
- relrae/relrae_components/config/RuBREx_conf.txt +2 -0
- relrae/relrae_components/config/pipeline_conf.txt +6 -0
- relrae/relrae_components/rules/SXPF_General_1_0.yaml +199 -0
- relrae/utils.py +148 -0
- relrae-0.9.0.dist-info/METADATA +122 -0
- relrae-0.9.0.dist-info/RECORD +20 -0
- relrae-0.9.0.dist-info/WHEEL +5 -0
- relrae-0.9.0.dist-info/entry_points.txt +2 -0
- relrae-0.9.0.dist-info/licenses/LICENSE +21 -0
- relrae-0.9.0.dist-info/top_level.txt +1 -0
relrae/RELRaE.py
ADDED
|
@@ -0,0 +1,91 @@
|
|
|
1
|
+
from rdflib import Graph, Namespace
|
|
2
|
+
from rdflib.namespace import RDF, RDFS, OWL
|
|
3
|
+
import xmlschema
|
|
4
|
+
from pathlib import Path
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
from .modules.RuBREx import RuBREx
|
|
8
|
+
from .modules.LLM_Ref import LLMRefinement
|
|
9
|
+
from .modules.human_fix import HumanFix
|
|
10
|
+
from .utils import get_now
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class RELRaE:
|
|
14
|
+
|
|
15
|
+
def __init__(self, onto_name, schema, namespace, prefix, config,
|
|
16
|
+
components_root="relrae_components", output_root="output"):
|
|
17
|
+
self.components_root = Path(components_root)
|
|
18
|
+
self.output_root = Path(output_root)
|
|
19
|
+
schema_path = self.components_root / "schema" / schema
|
|
20
|
+
self.schema = xmlschema.XMLSchema(schema_path)
|
|
21
|
+
self.namespace = Namespace(namespace)
|
|
22
|
+
self.prefix = prefix
|
|
23
|
+
self.onto_name = onto_name
|
|
24
|
+
self.onto = Graph()
|
|
25
|
+
self.errors = []
|
|
26
|
+
self.config = config
|
|
27
|
+
self.configs = [{"Pipeline": config}]
|
|
28
|
+
self.logs = {}
|
|
29
|
+
|
|
30
|
+
self.onto.bind('rdf', RDF)
|
|
31
|
+
self.onto.bind('rdfs', RDFS)
|
|
32
|
+
self.onto.bind('owl', OWL)
|
|
33
|
+
self.onto.bind(self.prefix, self.namespace)
|
|
34
|
+
|
|
35
|
+
def RuBREx(self):
|
|
36
|
+
m_rubrex = RuBREx(
|
|
37
|
+
self.schema, self.onto, self.prefix, self.namespace,
|
|
38
|
+
self.components_root,
|
|
39
|
+
)
|
|
40
|
+
m_rubrex.match_concepts()
|
|
41
|
+
m_rubrex.schema_coverage()
|
|
42
|
+
|
|
43
|
+
self.configs.append({"RuBREx": m_rubrex.config})
|
|
44
|
+
self.logs["RuBREx"] = m_rubrex.log
|
|
45
|
+
# NOTE: Consider the module complete
|
|
46
|
+
self.onto = self.onto + m_rubrex.onto
|
|
47
|
+
self.errors.append(m_rubrex.errors)
|
|
48
|
+
|
|
49
|
+
def LLM_Refinement_Loop(self):
|
|
50
|
+
m_LLM_ref = LLMRefinement(
|
|
51
|
+
self.schema, self.onto, self.prefix, self.namespace,
|
|
52
|
+
self.components_root)
|
|
53
|
+
m_LLM_ref.evaluate_relations()
|
|
54
|
+
self.errors.append(m_LLM_ref.errors)
|
|
55
|
+
|
|
56
|
+
def human_fix(self):
|
|
57
|
+
# TODO:
|
|
58
|
+
m_human_fix = HumanFix(
|
|
59
|
+
self.schema, self.onto, self.prefix, self.namespace,
|
|
60
|
+
self.components_root, self.errors, self.config["modules"]
|
|
61
|
+
)
|
|
62
|
+
m_human_fix.set_user_info()
|
|
63
|
+
m_human_fix.fix_errors()
|
|
64
|
+
m_human_fix.concepts_to_add_output(path)
|
|
65
|
+
|
|
66
|
+
def write_logs(self, path):
|
|
67
|
+
for key, values in self.logs.items():
|
|
68
|
+
with (path / f"{key}.txt").open("w", encoding="utf-8") as f:
|
|
69
|
+
for line in values:
|
|
70
|
+
f.write(line + "\n")
|
|
71
|
+
|
|
72
|
+
def write_metadata(self, path):
|
|
73
|
+
with (path / "metadata.txt").open("w", encoding="utf-8") as f:
|
|
74
|
+
for config in self.configs:
|
|
75
|
+
for section, values in config.items():
|
|
76
|
+
f.write(f"{section}\n\n")
|
|
77
|
+
for key, value in values.items():
|
|
78
|
+
f.write(f"{key}: {value}\n")
|
|
79
|
+
f.write("\n\n")
|
|
80
|
+
|
|
81
|
+
def serialise(self):
|
|
82
|
+
timestamp = get_now().replace(":", "-")
|
|
83
|
+
main_path = self.output_root / f"{self.onto_name}{timestamp}"
|
|
84
|
+
log_path = main_path / "logs"
|
|
85
|
+
log_path.mkdir(parents=True)
|
|
86
|
+
self.write_logs(log_path)
|
|
87
|
+
self.write_metadata(main_path)
|
|
88
|
+
self.onto.serialize(
|
|
89
|
+
destination=main_path / f"{self.onto_name}.ttl",
|
|
90
|
+
format="ttl",
|
|
91
|
+
)
|
relrae/__init__.py
ADDED
relrae/__main__.py
ADDED
relrae/_version.py
ADDED
relrae/cli.py
ADDED
|
@@ -0,0 +1,65 @@
|
|
|
1
|
+
import argparse
|
|
2
|
+
from importlib.resources import files
|
|
3
|
+
from pathlib import Path
|
|
4
|
+
import shutil
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
COMPONENT_DIRECTORIES = ("config", "rules", "schema", "output")
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def _copy_defaults(source, destination):
|
|
11
|
+
"""Copy packaged defaults without replacing user-edited files."""
|
|
12
|
+
for resource in source.iterdir():
|
|
13
|
+
# Finder metadata is not part of a RELRaE project, even if it happens
|
|
14
|
+
# to be present in a source checkout.
|
|
15
|
+
if resource.name.startswith("."):
|
|
16
|
+
continue
|
|
17
|
+
|
|
18
|
+
target = destination / resource.name
|
|
19
|
+
if resource.is_dir():
|
|
20
|
+
target.mkdir(exist_ok=True)
|
|
21
|
+
_copy_defaults(resource, target)
|
|
22
|
+
elif not target.exists():
|
|
23
|
+
with resource.open("rb") as source_file, target.open("xb") as target_file:
|
|
24
|
+
shutil.copyfileobj(source_file, target_file)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def initialise_project(project_root=None):
|
|
28
|
+
"""Create a RELRaE project beneath *project_root* (the cwd by default).
|
|
29
|
+
|
|
30
|
+
Existing files are intentionally left untouched, making ``relrae init``
|
|
31
|
+
safe to run again after configuration files have been customised.
|
|
32
|
+
"""
|
|
33
|
+
project_root = Path.cwd() if project_root is None else Path(project_root)
|
|
34
|
+
components_root = project_root / "relrae_components"
|
|
35
|
+
|
|
36
|
+
project_root.mkdir(parents=True, exist_ok=True)
|
|
37
|
+
for directory in COMPONENT_DIRECTORIES:
|
|
38
|
+
(components_root / directory).mkdir(parents=True, exist_ok=True)
|
|
39
|
+
|
|
40
|
+
packaged_defaults = files("relrae").joinpath("relrae_components")
|
|
41
|
+
_copy_defaults(packaged_defaults, components_root)
|
|
42
|
+
|
|
43
|
+
print(f"Initialised RELRaE project in {components_root}")
|
|
44
|
+
return components_root
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def main():
|
|
48
|
+
parser = argparse.ArgumentParser()
|
|
49
|
+
subparsers = parser.add_subparsers(dest="command")
|
|
50
|
+
|
|
51
|
+
init_parser = subparsers.add_parser(
|
|
52
|
+
"init", help="create a RELRaE project in the current directory"
|
|
53
|
+
)
|
|
54
|
+
init_parser.add_argument(
|
|
55
|
+
"directory",
|
|
56
|
+
nargs="?",
|
|
57
|
+
default=Path.cwd(),
|
|
58
|
+
type=Path,
|
|
59
|
+
help="project directory (default: current directory)",
|
|
60
|
+
)
|
|
61
|
+
|
|
62
|
+
args = parser.parse_args()
|
|
63
|
+
|
|
64
|
+
if args.command == "init":
|
|
65
|
+
initialise_project(args.directory)
|
|
@@ -0,0 +1,525 @@
|
|
|
1
|
+
import json
|
|
2
|
+
import os
|
|
3
|
+
import re
|
|
4
|
+
import configparser
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
from urllib import request
|
|
7
|
+
from urllib.error import URLError
|
|
8
|
+
from typing import Literal
|
|
9
|
+
from pydantic import BaseModel, Field, config
|
|
10
|
+
from ..utils import get_now, namespace_to_prefix
|
|
11
|
+
from random import randint
|
|
12
|
+
import numpy as np
|
|
13
|
+
from rdflib.namespace import RDFS
|
|
14
|
+
from rdflib import URIRef
|
|
15
|
+
import rdflib
|
|
16
|
+
|
|
17
|
+
class LLMRefinement:
|
|
18
|
+
|
|
19
|
+
def __init__(self, schema, onto, prefix, namespace,
|
|
20
|
+
components_root="relrae_components"):
|
|
21
|
+
self.log = []
|
|
22
|
+
cfg = configparser.ConfigParser()
|
|
23
|
+
cfg.read(Path(components_root) / "config" / "LLM_ref_conf.txt")
|
|
24
|
+
self.config = cfg["MAIN"]
|
|
25
|
+
self.errors = []
|
|
26
|
+
self.schema = schema
|
|
27
|
+
self.onto = onto
|
|
28
|
+
self.prefix = prefix
|
|
29
|
+
self.namespace = namespace
|
|
30
|
+
self.set_models()
|
|
31
|
+
|
|
32
|
+
# NOTE: This might not be neccessary
|
|
33
|
+
def set_models(self):
|
|
34
|
+
pass
|
|
35
|
+
|
|
36
|
+
def get_relations(self):
|
|
37
|
+
query = """
|
|
38
|
+
PREFIX rdf: <http://www.w3.org/1999/02/22-rdf-syntax-ns#>
|
|
39
|
+
PREFIX owl: <http://www.w3.org/2002/07/owl#>
|
|
40
|
+
|
|
41
|
+
SELECT ?property ?propertyType
|
|
42
|
+
WHERE {
|
|
43
|
+
VALUES ?propertyType {
|
|
44
|
+
rdf:Property
|
|
45
|
+
owl:ObjectProperty
|
|
46
|
+
owl:DatatypeProperty
|
|
47
|
+
owl:AnnotationProperty
|
|
48
|
+
}
|
|
49
|
+
|
|
50
|
+
?property rdf:type ?propertyType .
|
|
51
|
+
}
|
|
52
|
+
"""
|
|
53
|
+
qres = self.onto.query(query)
|
|
54
|
+
property_list = []
|
|
55
|
+
for r in qres:
|
|
56
|
+
property_list.append(r.property)
|
|
57
|
+
self.log.append(f"Retrieved {len(property_list)} properties")
|
|
58
|
+
return property_list
|
|
59
|
+
|
|
60
|
+
def get_relation_info(self, relation):
|
|
61
|
+
query = f"""
|
|
62
|
+
PREFIX rdf: <http://www.w3.org/1999/02/22-rdf-syntax-ns#>
|
|
63
|
+
PREFIX rdfs: <http://www.w3.org/2000/01/rdf-schema#>
|
|
64
|
+
PREFIX owl: <http://www.w3.org/2002/07/owl#>
|
|
65
|
+
|
|
66
|
+
SELECT ?subject ?predicate ?object
|
|
67
|
+
WHERE {{
|
|
68
|
+
{{
|
|
69
|
+
BIND(<{relation}> AS ?subject)
|
|
70
|
+
?subject ?predicate ?object .
|
|
71
|
+
}}
|
|
72
|
+
UNION
|
|
73
|
+
{{
|
|
74
|
+
?subject ?predicate <{relation}> .
|
|
75
|
+
BIND(<{relation}> AS ?object)
|
|
76
|
+
}}
|
|
77
|
+
}}
|
|
78
|
+
"""
|
|
79
|
+
qres = self.onto.query(query)
|
|
80
|
+
prefixed_relation = namespace_to_prefix(
|
|
81
|
+
relation, self.namespace, self.prefix)
|
|
82
|
+
rel_info = {prefixed_relation: []}
|
|
83
|
+
for r in qres:
|
|
84
|
+
rel_info[prefixed_relation].append([
|
|
85
|
+
namespace_to_prefix(r.predicate, self.namespace, self.prefix),
|
|
86
|
+
namespace_to_prefix(r.object, self.namespace, self.prefix)
|
|
87
|
+
])
|
|
88
|
+
self.log.append(f"Retrieved info for relation {str(relation)}")
|
|
89
|
+
# print(rel_info)
|
|
90
|
+
return rel_info
|
|
91
|
+
|
|
92
|
+
def normalise_relation_label(self, label):
|
|
93
|
+
label = namespace_to_prefix(label, self.namespace, self.prefix).strip()
|
|
94
|
+
prefix_marker = f"{self.prefix}:"
|
|
95
|
+
|
|
96
|
+
if ":" in label:
|
|
97
|
+
prefix, label = label.split(":", 1)
|
|
98
|
+
if prefix != self.prefix:
|
|
99
|
+
prefix = self.prefix
|
|
100
|
+
else:
|
|
101
|
+
prefix = self.prefix
|
|
102
|
+
|
|
103
|
+
words = re.findall(r"[A-Za-z0-9]+", label)
|
|
104
|
+
if not words:
|
|
105
|
+
raise ValueError(f"Cannot normalise empty relation label: {label}")
|
|
106
|
+
|
|
107
|
+
camel_label = words[0].lower()
|
|
108
|
+
for word in words[1:]:
|
|
109
|
+
camel_label += word[:1].upper() + word[1:].lower()
|
|
110
|
+
|
|
111
|
+
return f"{prefix_marker}{camel_label}"
|
|
112
|
+
|
|
113
|
+
def get_examples(self, model_role, strat):
|
|
114
|
+
if model_role == "eval":
|
|
115
|
+
examples = self.config["eval_examples"]
|
|
116
|
+
else:
|
|
117
|
+
examples = self.config["ref_examples"]
|
|
118
|
+
|
|
119
|
+
# print(examples)
|
|
120
|
+
shots = {"zero": None,
|
|
121
|
+
"one": 0,
|
|
122
|
+
"few": 4}
|
|
123
|
+
|
|
124
|
+
if strat == "zero":
|
|
125
|
+
return []
|
|
126
|
+
|
|
127
|
+
examples_list = examples[0:shots[strat]]
|
|
128
|
+
return examples_list
|
|
129
|
+
|
|
130
|
+
def average_eval(self, responses):
|
|
131
|
+
evals = []
|
|
132
|
+
confs = []
|
|
133
|
+
for r in responses:
|
|
134
|
+
if r["evaluation"] == "Yes":
|
|
135
|
+
evals.append(1)
|
|
136
|
+
else:
|
|
137
|
+
evals.append(0)
|
|
138
|
+
confs.append(r["confidence"])
|
|
139
|
+
mean_evals = round((sum(evals)/len(evals)), 0)
|
|
140
|
+
if int(mean_evals) == 1:
|
|
141
|
+
eval = "Yes"
|
|
142
|
+
else:
|
|
143
|
+
eval = "No"
|
|
144
|
+
avg_eval = {
|
|
145
|
+
"evaluation": eval,
|
|
146
|
+
"confidence": round(sum(confs)/len(confs), 3),
|
|
147
|
+
"variance": round(float(np.var(evals)), 3)
|
|
148
|
+
}
|
|
149
|
+
return avg_eval
|
|
150
|
+
|
|
151
|
+
def evaluation_loop(self, relation, evaluator, refiner, info):
|
|
152
|
+
accepted = False
|
|
153
|
+
rejected_labels = []
|
|
154
|
+
loops = 0
|
|
155
|
+
active_label = relation
|
|
156
|
+
|
|
157
|
+
while not accepted and loops < int(self.config["refinement_loops"]):
|
|
158
|
+
loops += 1
|
|
159
|
+
eval_model = evaluator[0]
|
|
160
|
+
eval_prompt = eval_model.build_eval_prompt(
|
|
161
|
+
namespace_to_prefix(active_label, self.namespace, self.prefix),
|
|
162
|
+
self.config["domain"],
|
|
163
|
+
self.config["source"],
|
|
164
|
+
info,
|
|
165
|
+
rejected_labels,
|
|
166
|
+
self.config["eval_messages"],
|
|
167
|
+
evaluator[1])
|
|
168
|
+
# print(eval_prompt)
|
|
169
|
+
full_eval = eval_model.run_prompt(eval_prompt)
|
|
170
|
+
if full_eval[0] == "Error":
|
|
171
|
+
self.log.append("Error: LLM could not generate valid response")
|
|
172
|
+
else:
|
|
173
|
+
avg_eval = self.average_eval(full_eval[0])
|
|
174
|
+
print(active_label)
|
|
175
|
+
print("\n=======================\n")
|
|
176
|
+
print(avg_eval)
|
|
177
|
+
print("\n=======================\n")
|
|
178
|
+
print(full_eval[1])
|
|
179
|
+
|
|
180
|
+
self.log.append(full_eval[1])
|
|
181
|
+
if avg_eval["evaluation"] == "Yes":
|
|
182
|
+
self.log.append("Label accepted")
|
|
183
|
+
break
|
|
184
|
+
|
|
185
|
+
self.log.append("Label rejected")
|
|
186
|
+
rejected_labels.append(namespace_to_prefix(
|
|
187
|
+
active_label, self.namespace, self.prefix))
|
|
188
|
+
print("\n ==== Rejected -> Refining ==== \n")
|
|
189
|
+
ref_model = refiner[0]
|
|
190
|
+
ref_prompt = ref_model.build_ref_prompt(
|
|
191
|
+
namespace_to_prefix(active_label, self.namespace, self.prefix),
|
|
192
|
+
self.config["domain"],
|
|
193
|
+
self.config["source"],
|
|
194
|
+
info,
|
|
195
|
+
rejected_labels,
|
|
196
|
+
self.config["ref_messages"],
|
|
197
|
+
refiner[1])
|
|
198
|
+
full_ref = ref_model.run_prompt(ref_prompt)
|
|
199
|
+
ref = full_ref[0][0]
|
|
200
|
+
print(ref)
|
|
201
|
+
print(rejected_labels)
|
|
202
|
+
print("\n=======================\n")
|
|
203
|
+
print(full_ref[1])
|
|
204
|
+
print("\n=======================\n")
|
|
205
|
+
self.log.append(full_ref[1])
|
|
206
|
+
active_label = self.normalise_relation_label(
|
|
207
|
+
ref["relationship_label"])
|
|
208
|
+
|
|
209
|
+
if accepted:
|
|
210
|
+
self.replace_relation(relation, active_label)
|
|
211
|
+
return ["refined", rejected_labels]
|
|
212
|
+
else:
|
|
213
|
+
return ["unrefined", rejected_labels]
|
|
214
|
+
|
|
215
|
+
def evaluate_relations(self):
|
|
216
|
+
relations = self.get_relations()
|
|
217
|
+
|
|
218
|
+
eval_m = LLM(self.config["eval_llm"],
|
|
219
|
+
EvaluatorResponse,
|
|
220
|
+
self.config["eval_repeats"],
|
|
221
|
+
self.namespace,
|
|
222
|
+
self.prefix)
|
|
223
|
+
eval_ex = self.get_examples("eval", self.config["eval_prompt_strat"])
|
|
224
|
+
|
|
225
|
+
ref_m = LLM(self.config["ref_llm"],
|
|
226
|
+
RefinerResponse,
|
|
227
|
+
self.config["ref_repeats"],
|
|
228
|
+
self.namespace,
|
|
229
|
+
self.prefix)
|
|
230
|
+
ref_ex = self.get_examples("ref", self.config["ref_prompt_strat"])
|
|
231
|
+
|
|
232
|
+
for rel in relations:
|
|
233
|
+
rel_info = self.get_relation_info(rel)
|
|
234
|
+
response = self.evaluation_loop(rel, [eval_m, eval_ex], [
|
|
235
|
+
ref_m, ref_ex], rel_info)
|
|
236
|
+
if response[0] == "unrefined":
|
|
237
|
+
self.errors.append([rel, rel_info, response[1]])
|
|
238
|
+
|
|
239
|
+
def replace_relation(self, original, new):
|
|
240
|
+
query = f"""
|
|
241
|
+
PREFIX rdf: <http://www.w3.org/1999/02/22-rdf-syntax-ns#>
|
|
242
|
+
PREFIX rdfs: <http://www.w3.org/2000/01/rdf-schema#>
|
|
243
|
+
PREFIX owl: <http://www.w3.org/2002/07/owl#>
|
|
244
|
+
|
|
245
|
+
SELECT ?s ?p ?o WHERE{{
|
|
246
|
+
?s ?p ?o
|
|
247
|
+
FITLER(
|
|
248
|
+
?s = <{self.namespace}{original}> ||
|
|
249
|
+
?p = <{self.namespace}{original}> ||
|
|
250
|
+
?o = <{self.namespace}{original}>
|
|
251
|
+
)
|
|
252
|
+
}}
|
|
253
|
+
"""
|
|
254
|
+
qres = self.onto.query(query)
|
|
255
|
+
|
|
256
|
+
for r in qres:
|
|
257
|
+
|
|
258
|
+
if r.p == str(RDFS.label):
|
|
259
|
+
n_sub = URIRef(f"{self.namespace}{new}")
|
|
260
|
+
n_obj = rdflib.Literal(new)
|
|
261
|
+
self.onto.remove((r.s, r.p, r.o))
|
|
262
|
+
self.onto.add((n_sub, r.p, n_obj))
|
|
263
|
+
elif r.s == f"{self.namespace}{original}":
|
|
264
|
+
n_sub = URIRef(f"{self.namespace}{new}")
|
|
265
|
+
self.onto.remove((r.s, r.p, r.o))
|
|
266
|
+
self.onto.add((n_sub, r.p, r.o))
|
|
267
|
+
elif r.p == f"{self.namespace}{original}":
|
|
268
|
+
n_pred = URIRef(f"{self.namespace}{new}")
|
|
269
|
+
self.onto.remove((r.s, r.p, r.o))
|
|
270
|
+
self.onto.add((r.s, n_pred, r.o))
|
|
271
|
+
elif r.o == f"{self.namespace}{original}":
|
|
272
|
+
n_obj = URIRef(f"{self.namespace}{new}")
|
|
273
|
+
self.onto.remove((r.s, r.p, r.o))
|
|
274
|
+
self.onto.add((r.s, r.p, n_obj))
|
|
275
|
+
|
|
276
|
+
s = URIRef(f"{self.namespace}{new}")
|
|
277
|
+
p = URIRef(f"{self.namespace}editedBy")
|
|
278
|
+
o = rdflib.Literal(f"{self.config["eval_llm"]} @ {get_now()}")
|
|
279
|
+
self.onto.add([s, p, o])
|
|
280
|
+
self.log.append(f"Relationship {original} updated to {new}.")
|
|
281
|
+
|
|
282
|
+
class EvaluatorResponse(BaseModel):
|
|
283
|
+
evaluation: Literal["Yes", "No"]
|
|
284
|
+
justification: str
|
|
285
|
+
confidence: float = Field(ge=0.00, le=100.00)
|
|
286
|
+
|
|
287
|
+
|
|
288
|
+
class RefinerResponse(BaseModel):
|
|
289
|
+
relationship_label: str
|
|
290
|
+
justification: str
|
|
291
|
+
|
|
292
|
+
|
|
293
|
+
class LLM:
|
|
294
|
+
|
|
295
|
+
def __init__(self, settings, format, repeats, ns, pr):
|
|
296
|
+
self.repeats = repeats
|
|
297
|
+
self.model = settings[0]
|
|
298
|
+
self.api_key = settings[1]
|
|
299
|
+
self.params = settings[2]
|
|
300
|
+
self.provider = self.get_provider(settings)
|
|
301
|
+
self.r_format = format
|
|
302
|
+
self.onto_ns = ns
|
|
303
|
+
self.onto_pr = pr
|
|
304
|
+
|
|
305
|
+
def gen_seeds(self, n):
|
|
306
|
+
seeds = []
|
|
307
|
+
for i in range(int(n)):
|
|
308
|
+
seeds.append(randint(1, 9999))
|
|
309
|
+
return seeds
|
|
310
|
+
|
|
311
|
+
def get_provider(self, settings):
|
|
312
|
+
if len(settings) > 3:
|
|
313
|
+
return str(settings[3]).lower()
|
|
314
|
+
if isinstance(self.params, dict) and self.params.get("provider"):
|
|
315
|
+
return str(self.params["provider"]).lower()
|
|
316
|
+
if str(self.api_key).lower() in ("", "none", "ollama"):
|
|
317
|
+
return "ollama"
|
|
318
|
+
if str(self.model).lower().startswith(("gpt", "o1", "o3", "o4")):
|
|
319
|
+
return "openai"
|
|
320
|
+
if str(self.model).lower().startswith(("gemini", "models/gemini")):
|
|
321
|
+
return "google"
|
|
322
|
+
return "ollama"
|
|
323
|
+
|
|
324
|
+
def build_messages(self, prompt, examples, active):
|
|
325
|
+
message_list = [{"role": "system", "content": prompt}]
|
|
326
|
+
for example in examples:
|
|
327
|
+
for e in example:
|
|
328
|
+
message_list.append(e)
|
|
329
|
+
message_list.append(active)
|
|
330
|
+
return message_list
|
|
331
|
+
|
|
332
|
+
def process_context(self, context):
|
|
333
|
+
context_list = []
|
|
334
|
+
reject_relations = [
|
|
335
|
+
"http://www.w3.org/2000/01/rdf-schema#label",
|
|
336
|
+
f"{self.onto_pr}:generatedBy",
|
|
337
|
+
f"{self.onto_pr}:hasXSDSource"
|
|
338
|
+
]
|
|
339
|
+
for l in context.keys():
|
|
340
|
+
for r in context[l]:
|
|
341
|
+
if r[0] in reject_relations:
|
|
342
|
+
continue
|
|
343
|
+
triple = [l, r[0], r[1]]
|
|
344
|
+
context_list.append(triple)
|
|
345
|
+
print(context_list)
|
|
346
|
+
return context_list
|
|
347
|
+
|
|
348
|
+
def build_eval_prompt(self, relation, domain, source, info, rejected_label, prompt, examples):
|
|
349
|
+
active = {"role": "user", "content": f"relation: {relation}, context: {
|
|
350
|
+
self.process_context(info)}, domain: {domain}, source: {source}, rejected labels: {rejected_label}"}
|
|
351
|
+
|
|
352
|
+
messages = self.build_messages(prompt[0], examples, active)
|
|
353
|
+
return messages
|
|
354
|
+
|
|
355
|
+
def build_ref_prompt(self, relation, domain, source, info, rejected_label, prompt, examples):
|
|
356
|
+
active = {"role": "user", "content": f"relation: {relation}, context: {
|
|
357
|
+
self.process_context(info)}, domain: {domain}, source: {source}, rejected_labels: {rejected_label}"}
|
|
358
|
+
|
|
359
|
+
messages = self.build_messages(prompt[0], examples, active)
|
|
360
|
+
return messages
|
|
361
|
+
|
|
362
|
+
def run_prompt(self, messages):
|
|
363
|
+
results = []
|
|
364
|
+
logs = []
|
|
365
|
+
seeds = []
|
|
366
|
+
seed = 0
|
|
367
|
+
|
|
368
|
+
while seed < int(self.repeats):
|
|
369
|
+
retrys = 10
|
|
370
|
+
valid = False
|
|
371
|
+
while not valid and retrys > 0:
|
|
372
|
+
c_seed = randint(1, 9999)
|
|
373
|
+
try:
|
|
374
|
+
raw_response = self.call_provider(messages, c_seed)
|
|
375
|
+
parsed_response = self.parse_response(raw_response)
|
|
376
|
+
results.append(parsed_response)
|
|
377
|
+
logs.append({
|
|
378
|
+
"provider": self.provider,
|
|
379
|
+
"model": self.model,
|
|
380
|
+
"seed": c_seed,
|
|
381
|
+
"response": parsed_response,
|
|
382
|
+
})
|
|
383
|
+
valid = True
|
|
384
|
+
seeds.append(c_seed)
|
|
385
|
+
seed += 1
|
|
386
|
+
retrys -= 1
|
|
387
|
+
except Exception as exc:
|
|
388
|
+
logs.append({
|
|
389
|
+
"provider": self.provider,
|
|
390
|
+
"model": self.model,
|
|
391
|
+
"seed": c_seed,
|
|
392
|
+
"error": str(exc),
|
|
393
|
+
})
|
|
394
|
+
retrys -= 1
|
|
395
|
+
|
|
396
|
+
if not results:
|
|
397
|
+
results = "Error"
|
|
398
|
+
|
|
399
|
+
return [results, logs]
|
|
400
|
+
|
|
401
|
+
def call_provider(self, messages, seed):
|
|
402
|
+
if self.provider == "ollama":
|
|
403
|
+
return self.call_ollama(messages, seed)
|
|
404
|
+
if self.provider == "openai":
|
|
405
|
+
return self.call_openai(messages, seed)
|
|
406
|
+
if self.provider in ("google", "googleai", "gemini"):
|
|
407
|
+
return self.call_google(messages, seed)
|
|
408
|
+
raise ValueError(f"Unsupported LLM provider: {self.provider}")
|
|
409
|
+
|
|
410
|
+
def get_params(self, seed, provider=None):
|
|
411
|
+
params = {}
|
|
412
|
+
if isinstance(self.params, dict):
|
|
413
|
+
params.update(self.params)
|
|
414
|
+
params.pop("provider", None)
|
|
415
|
+
params["seed"] = seed
|
|
416
|
+
if provider == "openai":
|
|
417
|
+
supported = {
|
|
418
|
+
"temperature",
|
|
419
|
+
"top_p",
|
|
420
|
+
"max_tokens",
|
|
421
|
+
"max_completion_tokens",
|
|
422
|
+
"presence_penalty",
|
|
423
|
+
"frequency_penalty",
|
|
424
|
+
"seed",
|
|
425
|
+
"stop",
|
|
426
|
+
}
|
|
427
|
+
params = {
|
|
428
|
+
key: value for key, value in params.items()
|
|
429
|
+
if key in supported
|
|
430
|
+
}
|
|
431
|
+
return params
|
|
432
|
+
|
|
433
|
+
def get_api_key(self, env_var):
|
|
434
|
+
if str(self.api_key).lower() not in ("", "none"):
|
|
435
|
+
return self.api_key
|
|
436
|
+
return os.environ.get(env_var)
|
|
437
|
+
|
|
438
|
+
def parse_response(self, raw_response):
|
|
439
|
+
if isinstance(raw_response, self.r_format):
|
|
440
|
+
return raw_response.model_dump()
|
|
441
|
+
if isinstance(raw_response, dict):
|
|
442
|
+
return self.r_format.model_validate(raw_response).model_dump()
|
|
443
|
+
return self.r_format.model_validate_json(raw_response).model_dump()
|
|
444
|
+
|
|
445
|
+
def messages_to_text(self, messages):
|
|
446
|
+
prompt_parts = []
|
|
447
|
+
for message in messages:
|
|
448
|
+
prompt_parts.append(f"{message['role']}: {message['content']}")
|
|
449
|
+
return "\n\n".join(prompt_parts)
|
|
450
|
+
|
|
451
|
+
def call_ollama(self, messages, seed):
|
|
452
|
+
params = self.get_params(seed, "ollama")
|
|
453
|
+
payload = {
|
|
454
|
+
"model": self.model,
|
|
455
|
+
"messages": messages,
|
|
456
|
+
"stream": False,
|
|
457
|
+
"format": self.r_format.model_json_schema(),
|
|
458
|
+
"options": params,
|
|
459
|
+
}
|
|
460
|
+
data = json.dumps(payload).encode("utf-8")
|
|
461
|
+
req = request.Request(
|
|
462
|
+
"http://localhost:11434/api/chat",
|
|
463
|
+
data=data,
|
|
464
|
+
headers={"Content-Type": "application/json"},
|
|
465
|
+
method="POST",
|
|
466
|
+
)
|
|
467
|
+
|
|
468
|
+
try:
|
|
469
|
+
with request.urlopen(req, timeout=120) as response:
|
|
470
|
+
body = json.loads(response.read().decode("utf-8"))
|
|
471
|
+
except URLError as exc:
|
|
472
|
+
raise RuntimeError(
|
|
473
|
+
"Could not reach Ollama at localhost:11434") from exc
|
|
474
|
+
|
|
475
|
+
return body["message"]["content"]
|
|
476
|
+
|
|
477
|
+
def call_openai(self, messages, seed):
|
|
478
|
+
try:
|
|
479
|
+
from openai import OpenAI
|
|
480
|
+
except ImportError as exc:
|
|
481
|
+
raise ImportError(
|
|
482
|
+
"Install the openai package to use OpenAI models") from exc
|
|
483
|
+
|
|
484
|
+
api_key = self.get_api_key("OPENAI_API_KEY")
|
|
485
|
+
if not api_key:
|
|
486
|
+
raise ValueError("OpenAI API key missing")
|
|
487
|
+
|
|
488
|
+
params = self.get_params(seed, "openai")
|
|
489
|
+
client = OpenAI(api_key=api_key)
|
|
490
|
+
completion = client.chat.completions.parse(
|
|
491
|
+
model=self.model,
|
|
492
|
+
messages=messages,
|
|
493
|
+
response_format=self.r_format,
|
|
494
|
+
**params,
|
|
495
|
+
)
|
|
496
|
+
return completion.choices[0].message.parsed
|
|
497
|
+
|
|
498
|
+
def call_google(self, messages, seed):
|
|
499
|
+
try:
|
|
500
|
+
from google import genai
|
|
501
|
+
except ImportError as exc:
|
|
502
|
+
raise ImportError(
|
|
503
|
+
"Install the google-genai package to use Google AI models"
|
|
504
|
+
) from exc
|
|
505
|
+
|
|
506
|
+
api_key = self.get_api_key("GEMINI_API_KEY")
|
|
507
|
+
if not api_key:
|
|
508
|
+
raise ValueError("Google AI API key missing")
|
|
509
|
+
|
|
510
|
+
params = self.get_params(seed, "google")
|
|
511
|
+
client = genai.Client(api_key=api_key)
|
|
512
|
+
response = client.models.generate_content(
|
|
513
|
+
model=self.model,
|
|
514
|
+
contents=self.messages_to_text(messages),
|
|
515
|
+
config={
|
|
516
|
+
**params,
|
|
517
|
+
"response_format": {
|
|
518
|
+
"text": {
|
|
519
|
+
"mime_type": "application/json",
|
|
520
|
+
"schema": self.r_format.model_json_schema(),
|
|
521
|
+
}
|
|
522
|
+
},
|
|
523
|
+
},
|
|
524
|
+
)
|
|
525
|
+
return response.text
|