diff --git a/tests/test_ccocr_evaluator.py b/tests/test_ccocr_evaluator.py new file mode 100644 index 000000000..6d1891d1f --- /dev/null +++ b/tests/test_ccocr_evaluator.py @@ -0,0 +1,31 @@ +import unittest + +from vlmeval.dataset.utils.ccocr_evaluator.doc_parsing_evaluator import ( + CustomConfig, + ParsingEvaluator, + TableTree, +) + + +class TestCCOCREvaluator(unittest.TestCase): + def test_teds_cell_distance_supports_long_content(self): + content_length = 2140 + predicted = TableTree("td", 1, 1, list("a" * content_length)) + ground_truth = TableTree("td", 1, 1, list("a" * (content_length - 1) + "b")) + + score = CustomConfig().rename(predicted, ground_truth) + + self.assertAlmostEqual(score, 1 / content_length) + + def test_doc_distance_supports_long_content(self): + content_length = 2140 + predicted = "a" * content_length + ground_truth = "a" * (content_length - 1) + "b" + + score = ParsingEvaluator("doc_parsing").eval_doc({"sample": predicted}, {"sample": ground_truth}) + + self.assertAlmostEqual(score, 1 - 1 / content_length) + + +if __name__ == "__main__": + unittest.main() diff --git a/vlmeval/dataset/utils/ccocr_evaluator/doc_parsing_evaluator.py b/vlmeval/dataset/utils/ccocr_evaluator/doc_parsing_evaluator.py index dbcb30b69..cd95b03da 100644 --- a/vlmeval/dataset/utils/ccocr_evaluator/doc_parsing_evaluator.py +++ b/vlmeval/dataset/utils/ccocr_evaluator/doc_parsing_evaluator.py @@ -1,7 +1,7 @@ import re from collections import deque -import nltk +import editdistance from apted import APTED, Config from apted.helpers import Tree from tqdm import tqdm @@ -67,7 +67,7 @@ def rename(self, node1, node2): return 1.0 if node1.tag == "td": if node1.content or node2.content: - return nltk.edit_distance(node1.content, node2.content) / max(len(node1.content), len(node2.content)) + return editdistance.eval(node1.content, node2.content) / max(len(node1.content), len(node2.content)) return 0.0 @@ -197,7 +197,7 @@ def eval_doc(self, response_info, gt_info): pred = pred.replace(' ', '').replace('\n', '') gt = gt.replace(' ', '').replace('\n', '') - edit_dist = nltk.edit_distance(pred, gt) / max(len(pred), len(gt)) + edit_dist = editdistance.eval(pred, gt) / max(len(pred), len(gt)) results.append(1 - edit_dist) score = sum(results) / len(results) @@ -246,7 +246,7 @@ def eval_formula(self, response_info, gt_info, op_name='formula'): elif op_name == 'molecular': pred = pred.replace("\n", "").replace(" ", "").replace("", "").replace("", "") gt = gt.replace(" ", "") - edit_dist = nltk.edit_distance(pred, gt) / max(len(pred), len(gt)) + edit_dist = editdistance.eval(pred, gt) / max(len(pred), len(gt)) results.append(1 - edit_dist) score = sum(results) / len(results) return score