Skip to content

Commit df0e695

Browse files
Merge pull request #80 from adgianv/expand_evaluator_class
Addressing issue #65 - Expanded evaluator class: extract results as a Dataframe
2 parents 97da28e + 98e3950 commit df0e695

4 files changed

Lines changed: 166 additions & 6 deletions

File tree

‎pyproject.toml‎

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,6 @@
22
requires = ["setuptools", "setuptools-scm"]
33
build-backend = "setuptools.build_meta"
44

5-
65
[project]
76
name = "nervaluate"
87
version = "0.2.0"
@@ -15,11 +14,14 @@ readme = "README.md"
1514
requires-python = ">=3.8"
1615
keywords = ["named-entity-recognition", "ner", "evaluation-metrics", "partial-match-scoring", "nlp"]
1716
license = {text = "MIT License"}
18-
classifiers=[
17+
classifiers = [
1918
"Programming Language :: Python :: 3",
2019
"Operating System :: OS Independent"
2120
]
2221

22+
[project.dependencies]
23+
pandas = "==2.0.1"
24+
2325
[project.urls]
2426
"Homepage" = "https://github.com/MantisAI/nervaluate"
25-
"Bug Tracker" = "https://github.com/MantisAI/nervaluate/issues"
27+
"Bug Tracker" = "https://github.com/MantisAI/nervaluate/issues"

‎requirements_dev.txt‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,4 +5,5 @@ gitchangelog
55
mypy==1.3.0
66
pre-commit==3.3.1
77
pylint==2.17.4
8-
pytest==7.3.1
8+
pytest==7.3.1
9+
pandas==2.0.1

‎src/nervaluate/evaluate.py‎

Lines changed: 48 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,8 @@
11
import logging
22
from copy import deepcopy
3-
from typing import List, Dict, Union, Tuple, Optional
3+
import pandas as pd
4+
from typing import List, Dict, Union, Tuple, Optional, Any
5+
from collections import defaultdict
46

57
from .utils import conll_to_spans, find_overlap, list_to_spans
68

@@ -118,6 +120,51 @@ def evaluate(self) -> Tuple[Dict, Dict, Dict, Dict]:
118120
)
119121

120122
return self.results, self.evaluation_agg_entities_type, self.evaluation_indices, self.evaluation_agg_indices
123+
124+
# Helper method to flatten a nested dictionary
125+
def _flatten_dict(self, d: Dict[str, Any], parent_key: str = '', sep: str = '.') -> Dict[str, Any]:
126+
"""
127+
Flattens a nested dictionary.
128+
129+
Args:
130+
d (dict): The dictionary to flatten.
131+
parent_key (str): The base key string to prepend to each dictionary key.
132+
sep (str): The separator to use when combining keys.
133+
134+
Returns:
135+
dict: A flattened dictionary.
136+
"""
137+
items: List[Tuple[str, Any]] = []
138+
for k, v in d.items():
139+
new_key = f"{parent_key}{sep}{k}" if parent_key else k
140+
if isinstance(v, dict):
141+
items.extend(self._flatten_dict(v, new_key, sep=sep).items())
142+
else:
143+
items.append((new_key, v))
144+
return dict(items)
145+
146+
# Modified results_to_dataframe method using the helper method
147+
def results_to_dataframe(self) -> Any:
148+
if not self.results:
149+
raise ValueError("self.results should be defined.")
150+
151+
if not isinstance(self.results, dict) or not all(isinstance(v, dict) for v in self.results.values()):
152+
raise ValueError("self.results must be a dictionary of dictionaries.")
153+
154+
# Flatten the nested results dictionary, including the 'entities' sub-dictionaries
155+
flattened_results: Dict[str, Dict[str, Any]] = {}
156+
for outer_key, inner_dict in self.results.items():
157+
flattened_inner_dict = self._flatten_dict(inner_dict)
158+
for inner_key, value in flattened_inner_dict.items():
159+
if inner_key not in flattened_results:
160+
flattened_results[inner_key] = {}
161+
flattened_results[inner_key][outer_key] = value
162+
163+
# Convert the flattened results to a pandas DataFrame
164+
try:
165+
return pd.DataFrame(flattened_results)
166+
except Exception as e:
167+
raise RuntimeError("Error converting flattened results to DataFrame") from e
121168

122169

123170
# flake8: noqa: C901

‎tests/test_evaluator.py‎

Lines changed: 111 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,116 @@
1-
# pylint: disable=C0302
1+
# pylint: disable=too-many-lines
2+
import pandas as pd
23
from nervaluate import Evaluator
34

5+
def test_results_to_dataframe():
6+
"""
7+
Test the results_to_dataframe method.
8+
"""
9+
# Setup
10+
evaluator = Evaluator(
11+
true=[['B-LOC', 'I-LOC', 'O'], ['B-PER', 'O', 'O']],
12+
pred=[['B-LOC', 'I-LOC', 'O'], ['B-PER', 'I-PER', 'O']],
13+
tags=['LOC', 'PER']
14+
)
15+
16+
# Mock results data for the purpose of this test
17+
evaluator.results = {
18+
'strict': {
19+
'correct': 10,
20+
'incorrect': 5,
21+
'partial': 3,
22+
'missed': 2,
23+
'spurious': 4,
24+
'precision': 0.625,
25+
'recall': 0.6667,
26+
'f1': 0.6452,
27+
'entities': {
28+
'LOC': {'correct': 4, 'incorrect': 1, 'partial': 0, 'missed': 1, 'spurious': 2},
29+
'PER': {'correct': 3, 'incorrect': 2, 'partial': 1, 'missed': 0, 'spurious': 1},
30+
'ORG': {'correct': 3, 'incorrect': 2, 'partial': 2, 'missed': 1, 'spurious': 1}
31+
}
32+
},
33+
'ent_type': {
34+
'correct': 8,
35+
'incorrect': 4,
36+
'partial': 1,
37+
'missed': 3,
38+
'spurious': 3,
39+
'precision': 0.5714,
40+
'recall': 0.6154,
41+
'f1': 0.5926,
42+
'entities': {
43+
'LOC': {'correct': 3, 'incorrect': 2, 'partial': 1, 'missed': 1, 'spurious': 1},
44+
'PER': {'correct': 2, 'incorrect': 1, 'partial': 0, 'missed': 2, 'spurious': 0},
45+
'ORG': {'correct': 3, 'incorrect': 1, 'partial': 0, 'missed': 0, 'spurious': 2}
46+
}
47+
},
48+
'partial': {
49+
'correct': 7,
50+
'incorrect': 3,
51+
'partial': 4,
52+
'missed': 1,
53+
'spurious': 5,
54+
'precision': 0.5385,
55+
'recall': 0.6364,
56+
'f1': 0.5833,
57+
'entities': {
58+
'LOC': {'correct': 2, 'incorrect': 1, 'partial': 1, 'missed': 1, 'spurious': 2},
59+
'PER': {'correct': 3, 'incorrect': 1, 'partial': 1, 'missed': 0, 'spurious': 1},
60+
'ORG': {'correct': 2, 'incorrect': 1, 'partial': 2, 'missed': 0, 'spurious': 2}
61+
}
62+
},
63+
'exact': {
64+
'correct': 9,
65+
'incorrect': 6,
66+
'partial': 2,
67+
'missed': 2,
68+
'spurious': 2,
69+
'precision': 0.6,
70+
'recall': 0.6429,
71+
'f1': 0.6207,
72+
'entities': {
73+
'LOC': {'correct': 4, 'incorrect': 1, 'partial': 0, 'missed': 1, 'spurious': 1},
74+
'PER': {'correct': 3, 'incorrect': 3, 'partial': 0, 'missed': 0, 'spurious': 0},
75+
'ORG': {'correct': 2, 'incorrect': 2, 'partial': 2, 'missed': 1, 'spurious': 1}
76+
}
77+
}
78+
}
79+
80+
# Expected DataFrame
81+
expected_data = {
82+
'correct': {'strict': 10, 'ent_type': 8, 'partial': 7, 'exact': 9},
83+
'incorrect': {'strict': 5, 'ent_type': 4, 'partial': 3, 'exact': 6},
84+
'partial': {'strict': 3, 'ent_type': 1, 'partial': 4, 'exact': 2},
85+
'missed': {'strict': 2, 'ent_type': 3, 'partial': 1, 'exact': 2},
86+
'spurious': {'strict': 4, 'ent_type': 3, 'partial': 5, 'exact': 2},
87+
'precision': {'strict': 0.625, 'ent_type': 0.5714, 'partial': 0.5385, 'exact': 0.6},
88+
'recall': {'strict': 0.6667, 'ent_type': 0.6154, 'partial': 0.6364, 'exact': 0.6429},
89+
'f1': {'strict': 0.6452, 'ent_type': 0.5926, 'partial': 0.5833, 'exact': 0.6207},
90+
'entities.LOC.correct': {'strict': 4, 'ent_type': 3, 'partial': 2, 'exact': 4},
91+
'entities.LOC.incorrect': {'strict': 1, 'ent_type': 2, 'partial': 1, 'exact': 1},
92+
'entities.LOC.partial': {'strict': 0, 'ent_type': 1, 'partial': 1, 'exact': 0},
93+
'entities.LOC.missed': {'strict': 1, 'ent_type': 1, 'partial': 1, 'exact': 1},
94+
'entities.LOC.spurious': {'strict': 2, 'ent_type': 1, 'partial': 2, 'exact': 1},
95+
'entities.PER.correct': {'strict': 3, 'ent_type': 2, 'partial': 3, 'exact': 3},
96+
'entities.PER.incorrect': {'strict': 2, 'ent_type': 1, 'partial': 1, 'exact': 3},
97+
'entities.PER.partial': {'strict': 1, 'ent_type': 0, 'partial': 1, 'exact': 0},
98+
'entities.PER.missed': {'strict': 0, 'ent_type': 2, 'partial': 0, 'exact': 0},
99+
'entities.PER.spurious': {'strict': 1, 'ent_type': 0, 'partial': 1, 'exact': 0},
100+
'entities.ORG.correct': {'strict': 3, 'ent_type': 3, 'partial': 2, 'exact': 2},
101+
'entities.ORG.incorrect': {'strict': 2, 'ent_type': 1, 'partial': 1, 'exact': 2},
102+
'entities.ORG.partial': {'strict': 2, 'ent_type': 0, 'partial': 2, 'exact': 2},
103+
'entities.ORG.missed': {'strict': 1, 'ent_type': 0, 'partial': 0, 'exact': 1},
104+
'entities.ORG.spurious': {'strict': 1, 'ent_type': 2, 'partial': 2, 'exact': 1}
105+
}
106+
107+
expected_df = pd.DataFrame(expected_data)
108+
109+
# Execute
110+
result_df = evaluator.results_to_dataframe()
111+
112+
# Assert
113+
pd.testing.assert_frame_equal(result_df, expected_df)
4114

5115
def test_evaluator_simple_case():
6116
true = [

0 commit comments

Comments
 (0)