graph-knowledge-doc-parser 0.1.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.
- graph_knowledge_doc_parser-0.1.0.dist-info/METADATA +326 -0
- graph_knowledge_doc_parser-0.1.0.dist-info/RECORD +38 -0
- graph_knowledge_doc_parser-0.1.0.dist-info/WHEEL +4 -0
- graph_knowledge_doc_parser-0.1.0.dist-info/entry_points.txt +3 -0
- kg_doc_parser/__init__.py +9 -0
- kg_doc_parser/cast_hinting.py +19 -0
- kg_doc_parser/document_ingester_logger.py +766 -0
- kg_doc_parser/models.py +277 -0
- kg_doc_parser/ocr.py +752 -0
- kg_doc_parser/pdf2png.py +286 -0
- kg_doc_parser/semantic_document_splitting_layerwise_edits.py +3302 -0
- kg_doc_parser/text_processing_utils.py +30 -0
- kg_doc_parser/utils/__init__.py +0 -0
- kg_doc_parser/utils/bounded_threadpool_executor.py +37 -0
- kg_doc_parser/utils/file_loaders.py +405 -0
- kg_doc_parser/utils/langchain.py +220 -0
- kg_doc_parser/utils/log.py +135 -0
- kg_doc_parser/utils/version_chaining.py +1278 -0
- kg_doc_parser/workflow_ingest/__init__.py +187 -0
- kg_doc_parser/workflow_ingest/_kogwistar.py +13 -0
- kg_doc_parser/workflow_ingest/adapters.py +212 -0
- kg_doc_parser/workflow_ingest/cache.py +63 -0
- kg_doc_parser/workflow_ingest/cli.py +324 -0
- kg_doc_parser/workflow_ingest/clients.py +444 -0
- kg_doc_parser/workflow_ingest/demo_harness.py +427 -0
- kg_doc_parser/workflow_ingest/design.py +208 -0
- kg_doc_parser/workflow_ingest/handlers.py +617 -0
- kg_doc_parser/workflow_ingest/models.py +575 -0
- kg_doc_parser/workflow_ingest/ocr_pipeline.py +1581 -0
- kg_doc_parser/workflow_ingest/page_index.py +473 -0
- kg_doc_parser/workflow_ingest/parser_core.py +862 -0
- kg_doc_parser/workflow_ingest/parsing.py +249 -0
- kg_doc_parser/workflow_ingest/probe.py +164 -0
- kg_doc_parser/workflow_ingest/providers.py +412 -0
- kg_doc_parser/workflow_ingest/runners.py +546 -0
- kg_doc_parser/workflow_ingest/semantics.py +231 -0
- kg_doc_parser/workflow_ingest/service.py +112 -0
- kg_doc_parser/workflow_ingest/smoke_assets.py +62 -0
kg_doc_parser/ocr.py
ADDED
|
@@ -0,0 +1,752 @@
|
|
|
1
|
+
if True:
|
|
2
|
+
import logging
|
|
3
|
+
import os
|
|
4
|
+
retry_failed_refine = False
|
|
5
|
+
logger = logging.getLogger(__name__)
|
|
6
|
+
logger.addHandler(logging.NullHandler())
|
|
7
|
+
ocr_json_version = "0.1"
|
|
8
|
+
import time
|
|
9
|
+
import base64
|
|
10
|
+
|
|
11
|
+
from langchain_core.callbacks import BaseCallbackHandler
|
|
12
|
+
from langchain_core.runnables import Runnable
|
|
13
|
+
from .models import NonText_box_2d, OCRClusterResponse, SplitPage, SplitPageMeta, NonTextCluster, TextCluster
|
|
14
|
+
from typing import Any, Iterable, cast, Callable, Optional, Literal, Union
|
|
15
|
+
try:
|
|
16
|
+
from typing import TypeAlias
|
|
17
|
+
except ImportError: # pragma: no cover
|
|
18
|
+
from typing_extensions import TypeAlias
|
|
19
|
+
import json
|
|
20
|
+
from pydantic_extension.model_slicing import (ModeSlicingMixin, NotMode, FrontendField, BackendField, LLMField,
|
|
21
|
+
DtoType,
|
|
22
|
+
BackendType,
|
|
23
|
+
FrontendType,
|
|
24
|
+
LLMType,
|
|
25
|
+
use_mode)
|
|
26
|
+
from pydantic_extension.model_slicing.mixin import ExcludeMode, DtoField
|
|
27
|
+
from pydantic import BaseModel, Field, model_validator, field_validator, field_serializer
|
|
28
|
+
from langchain_core.messages import SystemMessage, BaseMessage, HumanMessage
|
|
29
|
+
from langchain_core.language_models import BaseChatModel
|
|
30
|
+
try:
|
|
31
|
+
from .workflow_ingest.providers import WorkflowProviderSettings, build_chat_model
|
|
32
|
+
except ImportError: # pragma: no cover
|
|
33
|
+
from kg_doc_parser.pdf2png import RawFileLoader
|
|
34
|
+
from .pdf2png import RawFileLoader
|
|
35
|
+
PastCompatibleSplitPage: TypeAlias = SplitPage
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def _build_ocr_llm(model_name: str, *, callbacks=None):
|
|
39
|
+
settings = WorkflowProviderSettings.from_env()
|
|
40
|
+
spec = settings.ocr.model_copy(update={"model": model_name})
|
|
41
|
+
return build_chat_model(spec, callbacks=callbacks)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def get_page_json(folder_path, page_num):
|
|
45
|
+
with open(os.path.join(folder_path, 'page_'+str(page_num)+'.json'), 'r') as f:
|
|
46
|
+
file_json_raw = json.load(f)
|
|
47
|
+
return file_json_raw
|
|
48
|
+
def regen_page(file_json_raw, use_raw):
|
|
49
|
+
# add compatible to union if want to compatible with past models
|
|
50
|
+
"""regen from json returned by SplitPage.to_doc(), can be view as SplitPage.FromJson(filepath)"""
|
|
51
|
+
p = PastCompatibleSplitPage(**file_json_raw)
|
|
52
|
+
if use_raw:
|
|
53
|
+
return p.dump_supercede_parse()
|
|
54
|
+
try:
|
|
55
|
+
res = p.to_doc()
|
|
56
|
+
except:
|
|
57
|
+
raise
|
|
58
|
+
return res
|
|
59
|
+
def regen_doc(folder_path, use_raw = False):
|
|
60
|
+
pages_nums = sorted((int(i.rsplit(".json",1)[0].split("page_",1)[1]) for i in os.listdir(folder_path) if i.endswith('.json') and i.startswith("page_")))
|
|
61
|
+
pages = []
|
|
62
|
+
split_pages = []
|
|
63
|
+
for pn in pages_nums:
|
|
64
|
+
try:
|
|
65
|
+
pages.append(get_page_json(folder_path, pn))
|
|
66
|
+
split_pages.append(regen_page(pages[-1], use_raw = use_raw))
|
|
67
|
+
except Exception as e:
|
|
68
|
+
folder_path,pn
|
|
69
|
+
print(f'error at page {pn}')
|
|
70
|
+
print(f'in file {folder_path}')
|
|
71
|
+
logger.error(f'error at page {pn}')
|
|
72
|
+
logger.error(f'in file {folder_path}')
|
|
73
|
+
raise
|
|
74
|
+
|
|
75
|
+
# pages = map( partial(get_page_json, folder_path= folder_path), pages_nums)
|
|
76
|
+
# split_pages = map(regen_page, pages)
|
|
77
|
+
full_doc = list(split_pages)
|
|
78
|
+
return full_doc
|
|
79
|
+
|
|
80
|
+
class box_2d(BaseModel):
|
|
81
|
+
box_2d: list[int] = Field(description = 'box y min, x min, y max and x max')
|
|
82
|
+
label : str = Field(description = 'text in the box')
|
|
83
|
+
id: int = Field(description = 'id of the text box in the page, autoincrement from 0')
|
|
84
|
+
class RawOCRResponse(BaseModel):
|
|
85
|
+
"""id/ cluster number must be all unique, for example, one of the ocr boxes_2d used id='1',
|
|
86
|
+
the first image box id (cluster numebr) will be '2', the next signature will be '3' """
|
|
87
|
+
boxes_2d : list[box_2d] = Field(description = 'ocr text response, description of x min, y min, xmax and y max. Share id uniqueness with all non-ocr blocs')
|
|
88
|
+
non_text_objects: DtoType[list[NonText_box_2d]] = Field(description="the non-OCR object results. Share cluster number uniqueness with OCR texts in box2d. "
|
|
89
|
+
"For example, if box2d list takes id 1, 2, 4, 5, non_text_objects will take up 3, 6... etc")
|
|
90
|
+
is_empty_page: DtoType[Optional[bool]] = Field(default = False, description="true if the whole page is empty without recognisable text.")
|
|
91
|
+
printed_page_number: DtoType[Optional[str]] = Field(description='the page number identified from OCR texts, can be in form of roman numerals such as "i", "ii", "iii", "iv"...; '
|
|
92
|
+
'Arabic numeral such as 1, 2, 3... or letter such as "a", "b", "c"...\n'
|
|
93
|
+
'Sometimes the are surrounded by symbols such as "- 1 -", "- 2 -"'
|
|
94
|
+
r"Can be null/none if there is no page order assigned and printed and found in the scanned texts. Do not assign page number. Only use page number found.")
|
|
95
|
+
meaningful_ordering : DtoType[list[int]] = Field(description="The correct meaningful ordering of the identified text clusters. Must cover all OCR_text_clusters once and only once. ")
|
|
96
|
+
page_x_min : DtoType[float]=Field(description='the page x min in pixel coordinate. ')
|
|
97
|
+
page_x_max : DtoType[float]=Field(description='the page x max in pixel coordinate. ')
|
|
98
|
+
page_y_min : DtoType[float]=Field(description='the page y min in pixel coordinate. ')
|
|
99
|
+
page_y_max : DtoType[float]=Field(description='the page y max in pixel coordinate. ')
|
|
100
|
+
estimated_rotation_degrees : DtoType[float]=Field(description='the page estimated rotation degree using right hand rule. ')
|
|
101
|
+
incomplete_words_on_edge: DtoType[bool] = Field(description='If there is any text being incomplete due to the scan does not scan the edges properly. ')
|
|
102
|
+
incomplete_text: DtoType[bool] = Field(description='Any incomplete text')
|
|
103
|
+
data_loss_likelihood: DtoType[float] = Field(description='The likelihood (range from 0.0 to 1.0 inclusive) that the page has lost information by missing the scan data on the edges of the page.' )
|
|
104
|
+
scan_quality: DtoType[Literal['low', 'medium', 'high']] = Field(description='The image quality of the scan. All qualities exclude signatures. '
|
|
105
|
+
'"low", "medium" or "high". '
|
|
106
|
+
'low: text barely legible. medium: Legible with non smooth due to pixelation. high: texts are easily and highly identifiable. ' )
|
|
107
|
+
contains_table: DtoType[bool] = Field(description='Whether this page contains table. ')
|
|
108
|
+
# is_signature_page: DtoType[bool] = Field(description='Whether this page contains signature. Must agree with signature_blocks')
|
|
109
|
+
# signature_blocks : DtoType[list[SignatureInfo]] = Field(default = [],
|
|
110
|
+
# description="The text cluster that belongs to signatory/signature (if any). "
|
|
111
|
+
# "Indicate whether it is signed or unsigned signatory. "
|
|
112
|
+
# "Share id uniqueness with OCR text boxes_2d. ")
|
|
113
|
+
|
|
114
|
+
@model_validator(mode='after')
|
|
115
|
+
def check_cluster_meaningful_ordering_agreement(self):
|
|
116
|
+
assert bool(self.is_empty_page) ^ (len(self.boxes_2d) > 0), f"is_empty_page value {self.is_empty_page} disagree with OCR_text_clusters len={len(self.boxes_2d)}"
|
|
117
|
+
overlap_id = set(i.id for i in self.non_text_objects).union(set(i.id for i in self.boxes_2d))
|
|
118
|
+
if not len([i.id for i in (self.non_text_objects + self.boxes_2d)]) == len(set(i.id for i in self.non_text_objects + self.boxes_2d)):
|
|
119
|
+
raise ValueError(f"cluster number from non_text_objects block and ocr text blocks must be ALL distinct. overlap ids {overlap_id}")
|
|
120
|
+
try:
|
|
121
|
+
if not (len(self.meaningful_ordering) == len(set(self.meaningful_ordering))): # <= len(self.OCR_text_clusters)):
|
|
122
|
+
raise ValueError("meaningful_order must cover each text cluster at most once")
|
|
123
|
+
except Exception as e:
|
|
124
|
+
raise e
|
|
125
|
+
return self
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
class OCRMetaResponse(BaseModel):
|
|
129
|
+
"meatada of an OCR page"
|
|
130
|
+
printed_page_number: DtoType[Optional[str]] = Field(description='the page number identified from OCR texts, can be in form of roman numerals such as "i", "ii", "iii", "iv"...; '
|
|
131
|
+
'Arabic numeral such as 1, 2, 3... or letter such as "a", "b", "c"...\n'
|
|
132
|
+
'Sometimes the are surrounded by symbols such as "- 1 -", "- 2 -"'
|
|
133
|
+
r"Can be null/none if there is no page order assigned and printed and found in the scanned texts. Do not assign page number. Only use page number found.")
|
|
134
|
+
|
|
135
|
+
page_x_min : DtoType[float]=Field(description='the page x min in pixel coordinate. ')
|
|
136
|
+
page_x_max : DtoType[float]=Field(description='the page x max in pixel coordinate. ')
|
|
137
|
+
page_y_min : DtoType[float]=Field(description='the page y min in pixel coordinate. ')
|
|
138
|
+
page_y_max : DtoType[float]=Field(description='the page y max in pixel coordinate. ')
|
|
139
|
+
estimated_rotation_degrees : DtoType[float]=Field(description='the page estimated rotation degree using right hand rule. ')
|
|
140
|
+
incomplete_words_on_edge: DtoType[bool] = Field(description='If there is any text being incomplete due to the scan does not scan the edges properly. ')
|
|
141
|
+
incomplete_text: DtoType[bool] = Field(description='Any incomplete text')
|
|
142
|
+
data_loss_likelihood: DtoType[float] = Field(description='The likelihood (range from 0.0 to 1.0 inclusive) that the page has lost information by missing the scan data on the edges of the page.' )
|
|
143
|
+
scan_quality: DtoType[Literal['low', 'medium', 'high']] = Field(description='The image quality of the scan. All qualities exclude signatures. '
|
|
144
|
+
'"low", "medium" or "high". '
|
|
145
|
+
'low: text barely legible. medium: Legible with non smooth due to pixelation. high: texts are easily and highly identifiable. ' )
|
|
146
|
+
contains_table: DtoType[bool] = Field(description='Whether this page contains table. ')
|
|
147
|
+
class OCRClusterResponseMetaless(ModeSlicingMixin, BaseModel):
|
|
148
|
+
"response of OCR once meta is quickly skimmed/ determined in separate run"
|
|
149
|
+
OCR_text_clusters: DtoType[list[TextCluster]] = Field(description="the OCR text results.")
|
|
150
|
+
non_text_objects: DtoType[list[NonText_box_2d]] = Field(description="the non-OCR object results. Share cluster number uniqueness with OCR texts. ")
|
|
151
|
+
printed_page_number: DtoType[Optional[str]] = Field(description='the page number identified from OCR texts, can be in form of roman numerals such as "i", "ii", "iii", "iv"...; '
|
|
152
|
+
'Arabic numeral such as 1, 2, 3... or letter such as "a", "b", "c"...\n'
|
|
153
|
+
'Sometimes the are surrounded by symbols such as "- 1 -", "- 2 -"'
|
|
154
|
+
r"Can be null/none if there is no page order assigned and printed and found in the scanned texts. Do not assign page number. Only use page number found.")
|
|
155
|
+
meaningful_ordering : DtoType[list[int]] = Field(description="The correct meaningful ordering of the identified text clusters. Must cover all OCR_text_clusters once and only once. ")
|
|
156
|
+
class RawOCRResponseMetaless(ModeSlicingMixin, BaseModel):
|
|
157
|
+
boxes_2d : list[box_2d] = Field(description = 'description of x min, y min, xmax and y max')
|
|
158
|
+
non_text_objects: DtoType[list[NonText_box_2d]] = Field(description="the non-OCR object results. Share cluster number uniqueness with OCR texts. ")
|
|
159
|
+
printed_page_number: DtoType[Optional[str]] = Field("",description='the page number identified from OCR texts, can be in form of roman numerals such as "i", "ii", "iii", "iv"...; '
|
|
160
|
+
'Arabic numeral such as 1, 2, 3... or letter such as "a", "b", "c"...\n'
|
|
161
|
+
'Sometimes the are surrounded by symbols such as "- 1 -", "- 2 -"'
|
|
162
|
+
r"Can be null/none if there is no page order assigned and printed and found in the scanned texts. Do not assign page number. Only use page number found.")
|
|
163
|
+
meaningful_ordering : DtoType[list[int]] = Field(description="The correct meaningful ordering of the identified text clusters. Must cover all OCR_text_clusters once and only once. ")
|
|
164
|
+
|
|
165
|
+
def get_first_round_response(draft_responses, llm: BaseChatModel, model_name: str, cb: BaseCallbackHandler,
|
|
166
|
+
messages: list[BaseMessage], sys_message, img_message, usage_metadata) -> OCRClusterResponse | None:
|
|
167
|
+
|
|
168
|
+
chain = llm.with_structured_output(RawOCRResponse, include_raw = True)
|
|
169
|
+
before_parse: Runnable = chain.steps[0]
|
|
170
|
+
after_parse: Runnable = chain.steps[1]
|
|
171
|
+
raw_response = before_parse.invoke(messages, config={"callbacks": [cb]}
|
|
172
|
+
)
|
|
173
|
+
|
|
174
|
+
|
|
175
|
+
if hasattr(raw_response,"usage_metadata"):
|
|
176
|
+
usage_metadata.append(raw_response.usage_metadata)
|
|
177
|
+
else:
|
|
178
|
+
usage_metadata.append(None)
|
|
179
|
+
response_with_raw: dict[str, RawOCRResponse] = after_parse.invoke(raw_response)
|
|
180
|
+
response: RawOCRResponse | OCRClusterResponse | None
|
|
181
|
+
response1 : RawOCRResponse | None = response_with_raw.get('parsed')
|
|
182
|
+
parsing_error = response_with_raw.get('parsing_error')
|
|
183
|
+
if response1 is None:
|
|
184
|
+
|
|
185
|
+
try:
|
|
186
|
+
raw = response_with_raw.get('raw')
|
|
187
|
+
temp = json.loads(raw.content[0]['text'])
|
|
188
|
+
response1 = RawOCRResponse.model_validate(temp)
|
|
189
|
+
|
|
190
|
+
except:
|
|
191
|
+
pass
|
|
192
|
+
if response1 is not None:
|
|
193
|
+
response = RawOCRResponse_to_OCRClusterResponse(response1)
|
|
194
|
+
else:
|
|
195
|
+
response = response1
|
|
196
|
+
if (response is None) or parsing_error:
|
|
197
|
+
|
|
198
|
+
sys_message_2 = sys_message.model_copy(deep = True)
|
|
199
|
+
class OCRDraftResponse(BaseModel):
|
|
200
|
+
text: str = Field(description = "OCR identified text with layout")
|
|
201
|
+
ocr_draft_response = cast(
|
|
202
|
+
OCRDraftResponse | None,
|
|
203
|
+
llm.with_structured_output(OCRDraftResponse).invoke(messages),
|
|
204
|
+
)
|
|
205
|
+
if (ocr_draft_response is not None) and (ocr_draft_response.text is not None) and ocr_draft_response.text != "":
|
|
206
|
+
draft_responses[model_name] = ocr_draft_response.text
|
|
207
|
+
sys_message_2.content += ("If your internal OCR fails. Focus on table parsing mode because my error analysis modes often show that the failing OCR pages are usually highly complicated tables. "
|
|
208
|
+
f"Try to put in as much data as possible given all text found by simple OCR for your reference:```{ocr_draft_response.text}```" if ocr_draft_response else"")
|
|
209
|
+
response_with_raw = cast(dict[str, RawOCRResponse], llm.with_structured_output(RawOCRResponse, include_raw = True).invoke(
|
|
210
|
+
[sys_message, img_message]
|
|
211
|
+
))
|
|
212
|
+
raw_ocr_response : None | RawOCRResponse= None
|
|
213
|
+
if response_with_raw.get('parsed'):
|
|
214
|
+
raw_ocr_response = cast(RawOCRResponse, response_with_raw.get('parsed'))
|
|
215
|
+
else:
|
|
216
|
+
try:
|
|
217
|
+
raw_ocr_response = RawOCRResponse.model_validate(json.loads(response_with_raw.get('raw').content))
|
|
218
|
+
except:
|
|
219
|
+
pass
|
|
220
|
+
if raw_ocr_response is not None:
|
|
221
|
+
response = RawOCRResponse_to_OCRClusterResponse(raw_ocr_response)
|
|
222
|
+
_raw = response_with_raw.get('raw')
|
|
223
|
+
parsing_error = response_with_raw.get('parsing_error')
|
|
224
|
+
return response
|
|
225
|
+
def validate_response_mutate_inplace(response: OCRClusterResponse | None, response_dict: dict, image_file_path, model_name, page_file_name):
|
|
226
|
+
|
|
227
|
+
if response is None:
|
|
228
|
+
logger.error(f"LLM returned None as response, file name = {image_file_path}, {model_name=}")
|
|
229
|
+
raise(ValueError(f"LLM returned None as response, file name = {image_file_path}, {model_name=}"))
|
|
230
|
+
else:
|
|
231
|
+
if len(response.OCR_text_clusters) == 0 or len(''.join([c.text for c in response.OCR_text_clusters])) == 0:
|
|
232
|
+
logger.info(f'emptydoc by {model_name}')
|
|
233
|
+
if model_name != "gemini-2.5-pro": # only trust the verdict from newest advanced model if nothing detected.
|
|
234
|
+
raise Exception(f"Empty OCR Page error. No text detected at all by less advanced model {model_name}. "
|
|
235
|
+
"Application only trust empty response from advanced model 'gemini-2.5-pro'")
|
|
236
|
+
else:
|
|
237
|
+
pass
|
|
238
|
+
else:
|
|
239
|
+
|
|
240
|
+
pass
|
|
241
|
+
response_dict_local = response.model_dump()
|
|
242
|
+
response_dict_local['pdf_page_num'] = page_file_name.rsplit('.',1)[0].rsplit("_",1)[-1]
|
|
243
|
+
response_dict_local['metadata'] = {"ocr_model_name": model_name, "ocr_datetime" : time.time(), "ocr_json_version": str(ocr_json_version)}
|
|
244
|
+
response_dict_local['refined_version'] = None
|
|
245
|
+
sp= SplitPage.model_validate(response_dict_local)
|
|
246
|
+
if sp is None:
|
|
247
|
+
sp = SplitPage(**response_dict_local)
|
|
248
|
+
try:
|
|
249
|
+
sp.to_doc()
|
|
250
|
+
response_dict.update(response_dict_local)
|
|
251
|
+
ok = True
|
|
252
|
+
except Exception as e:
|
|
253
|
+
logger.error(f"Generated json fail to reproduce doc, file name = {image_file_path}, {model_name=}")
|
|
254
|
+
logger.error(e)
|
|
255
|
+
sp.to_doc()
|
|
256
|
+
raise(ValueError(f"Generated json fail to reproduce doc, file name = {image_file_path}, {model_name=}"))
|
|
257
|
+
return sp
|
|
258
|
+
class TextBox(BaseModel):
|
|
259
|
+
text: str = Field(description = 'identified text')
|
|
260
|
+
bounding_box : list[int] = Field(description = 'Bounding box of identified text')
|
|
261
|
+
id: int = Field(description = 'id of the text box in the page, autoincrement from 0')
|
|
262
|
+
class NonTextObject(BaseModel):
|
|
263
|
+
text: str = Field(description = 'identified non-OCR object')
|
|
264
|
+
bounding_box : list[int] = Field(description = 'Bounding box of identified non-OCR object')
|
|
265
|
+
id: int = Field(description = 'id of the text box in the page, autoincrement from 0')
|
|
266
|
+
class TextBoxResponse(BaseModel):
|
|
267
|
+
text_blocks: list[TextBox] = Field(description = "bounding boxes and text identified")
|
|
268
|
+
non_text_blocks: list[NonTextObject] = Field(description = "bounding boxes and description of the object identified")
|
|
269
|
+
meaningful_ordering : DtoType[list[int]] = Field(description="The correct meaningful ordering of the identified text clusters. Must cover all OCR_text_clusters once and only once. ")
|
|
270
|
+
printed_page_number: str = Field("",description='the page number identified')
|
|
271
|
+
def RawOCRResponse_to_OCRClusterResponse(raw_response: RawOCRResponse | RawOCRResponseMetaless | TextBoxResponse) -> OCRClusterResponse:
|
|
272
|
+
temp = raw_response.model_dump()
|
|
273
|
+
boxes_2d = temp.pop('boxes_2d') # y min x min, y max x max
|
|
274
|
+
temp['OCR_text_clusters'] = [TextCluster.model_validate({"text" : i['label'],
|
|
275
|
+
"bb_y_min" : i['box_2d'][0],
|
|
276
|
+
"bb_x_min" : i['box_2d'][1],
|
|
277
|
+
"bb_y_max" : i['box_2d'][2],
|
|
278
|
+
"bb_x_max" : i['box_2d'][3],
|
|
279
|
+
"cluster_number" : i['id']}) for i in boxes_2d]
|
|
280
|
+
non_text_objects = temp.pop('non_text_objects')
|
|
281
|
+
temp['non_text_objects'] = [NonTextCluster.model_validate({"description" : i.get('label', i['description']),
|
|
282
|
+
"bb_y_min" : i['box_2d'][0],
|
|
283
|
+
"bb_x_min" : i['box_2d'][1],
|
|
284
|
+
"bb_y_max" : i['box_2d'][2],
|
|
285
|
+
"bb_x_max" : i['box_2d'][3],
|
|
286
|
+
"cluster_number" : i['id']}) for i in non_text_objects]
|
|
287
|
+
return OCRClusterResponse.model_validate(temp)
|
|
288
|
+
def final_resort(draft_responses: dict, messages, page_file_name, model_name, image_file_path):
|
|
289
|
+
"""
|
|
290
|
+
One day gemini suddenly cannot run but return a totally different schema, ad hoc code fix to fit the transformed schema and
|
|
291
|
+
break down document reading into 2 tasks, namely meta and ocr and non ocr recognition
|
|
292
|
+
"""
|
|
293
|
+
max_v = ""
|
|
294
|
+
for k, v in draft_responses.items():
|
|
295
|
+
if len(v) > len(max_v):
|
|
296
|
+
max_k = k,
|
|
297
|
+
max_v = v
|
|
298
|
+
earlier_partial_ocr = draft_responses.get("gemini-2.5-pro") or draft_responses.get("gemini-2.5-flash") or max_v
|
|
299
|
+
llm = _build_ocr_llm("gemini-2.5-pro", callbacks=[cb])
|
|
300
|
+
|
|
301
|
+
ocr_meta_response: OCRMetaResponse| None =cast (OCRMetaResponse | None , llm.with_structured_output(OCRMetaResponse).invoke(messages[:2]))
|
|
302
|
+
if ocr_meta_response is None:
|
|
303
|
+
raise Exception("model capability cannot even skim meta coarse level information")
|
|
304
|
+
metadata = [SystemMessage("Earlier steps has already determined the metadata about this document: \n\n" + str(ocr_meta_response.model_dump()))]
|
|
305
|
+
has_error = False
|
|
306
|
+
try:
|
|
307
|
+
response2:RawOCRResponseMetaless | None= cast(RawOCRResponseMetaless | None,
|
|
308
|
+
llm.with_structured_output(RawOCRResponseMetaless).invoke(messages[:2] + metadata))
|
|
309
|
+
if response2 is None:
|
|
310
|
+
has_error = True
|
|
311
|
+
else:
|
|
312
|
+
response = RawOCRResponse_to_OCRClusterResponse(response2)
|
|
313
|
+
except Exception as _e:
|
|
314
|
+
has_error = True
|
|
315
|
+
if has_error:
|
|
316
|
+
# retry only
|
|
317
|
+
try:
|
|
318
|
+
response3:TextBoxResponse | None = cast (TextBoxResponse | None , llm.with_structured_output(TextBoxResponse).invoke(messages[:2]))
|
|
319
|
+
if response3 is None:
|
|
320
|
+
has_error = True
|
|
321
|
+
raise(Exception("error when trying TextBoxResponse"))
|
|
322
|
+
else:
|
|
323
|
+
response = RawOCRResponse_to_OCRClusterResponse(response3)
|
|
324
|
+
|
|
325
|
+
coered_response = OCRClusterResponse(
|
|
326
|
+
printed_page_number = ocr_meta_response.printed_page_number,
|
|
327
|
+
page_x_min = ocr_meta_response.page_x_min,
|
|
328
|
+
page_x_max = ocr_meta_response.page_x_max,
|
|
329
|
+
page_y_min = ocr_meta_response.page_y_min,
|
|
330
|
+
page_y_max = ocr_meta_response.page_y_max,
|
|
331
|
+
estimated_rotation_degrees = ocr_meta_response.estimated_rotation_degrees,
|
|
332
|
+
incomplete_words_on_edge = ocr_meta_response.incomplete_words_on_edge,
|
|
333
|
+
incomplete_text = ocr_meta_response.incomplete_text,
|
|
334
|
+
data_loss_likelihood = ocr_meta_response.data_loss_likelihood,
|
|
335
|
+
scan_quality = ocr_meta_response.scan_quality,
|
|
336
|
+
contains_table = ocr_meta_response.contains_table,
|
|
337
|
+
|
|
338
|
+
**response.model_dump()
|
|
339
|
+
)
|
|
340
|
+
except Exception as _e:
|
|
341
|
+
try:
|
|
342
|
+
coered_response = OCRClusterResponse(
|
|
343
|
+
OCR_text_clusters = [TextCluster(text = earlier_partial_ocr,
|
|
344
|
+
bb_x_min = ocr_meta_response.page_x_min,
|
|
345
|
+
bb_y_min = ocr_meta_response.page_y_min,
|
|
346
|
+
bb_x_max =ocr_meta_response.page_x_max,
|
|
347
|
+
bb_y_max =ocr_meta_response.page_y_max,
|
|
348
|
+
cluster_number = 0)],
|
|
349
|
+
non_text_objects=[],
|
|
350
|
+
meaningful_ordering = [0],
|
|
351
|
+
printed_page_number = ocr_meta_response.printed_page_number,
|
|
352
|
+
page_x_min = ocr_meta_response.page_x_min,
|
|
353
|
+
page_x_max = ocr_meta_response.page_x_max,
|
|
354
|
+
page_y_min = ocr_meta_response.page_y_min,
|
|
355
|
+
page_y_max = ocr_meta_response.page_y_max,
|
|
356
|
+
estimated_rotation_degrees = ocr_meta_response.estimated_rotation_degrees,
|
|
357
|
+
incomplete_words_on_edge = ocr_meta_response.incomplete_words_on_edge,
|
|
358
|
+
incomplete_text = ocr_meta_response.incomplete_text,
|
|
359
|
+
data_loss_likelihood = ocr_meta_response.data_loss_likelihood,
|
|
360
|
+
scan_quality = ocr_meta_response.scan_quality,
|
|
361
|
+
contains_table = ocr_meta_response.contains_table,
|
|
362
|
+
)
|
|
363
|
+
response_dict = coered_response.model_dump()
|
|
364
|
+
response_dict['pdf_page_num'] = page_file_name.rsplit('.',1)[0].rsplit("_",1)[-1]
|
|
365
|
+
response_dict['metadata'] = {"ocr_model_name": model_name, "ocr_datetime" : time.time(), "ocr_json_version": str(ocr_json_version)}
|
|
366
|
+
response_dict['refined_version'] = None
|
|
367
|
+
|
|
368
|
+
sp= SplitPage.model_validate(response_dict)
|
|
369
|
+
if sp is None:
|
|
370
|
+
sp = SplitPage(**response_dict)
|
|
371
|
+
try:
|
|
372
|
+
sp.to_doc()
|
|
373
|
+
ok = True
|
|
374
|
+
except Exception as _e:
|
|
375
|
+
raise Exception("Validation error response_dict cannot be validate into SplitPage")
|
|
376
|
+
except Exception as _e:
|
|
377
|
+
raise(ValueError(f"All LLM failed and coercing final resort fail, file name = {image_file_path}"))
|
|
378
|
+
def TextBoxResponsePlusMetaResponse_to_OCRClusterResponse(raw_response: TextBoxResponse, meta_response: OCRMetaResponse) -> OCRClusterResponse:
|
|
379
|
+
|
|
380
|
+
|
|
381
|
+
temp = meta_response.model_dump()
|
|
382
|
+
temp.update(raw_response.model_dump())
|
|
383
|
+
text_blocks = temp.pop('text_blocks') # y min x min, y max x max
|
|
384
|
+
temp['OCR_text_clusters'] = [TextCluster.model_validate({"text" : i['text'],
|
|
385
|
+
"bb_y_min" : i['bounding_box'][0],
|
|
386
|
+
"bb_x_min" : i['bounding_box'][1],
|
|
387
|
+
"bb_y_max" : i['bounding_box'][2],
|
|
388
|
+
"bb_x_max" : i['bounding_box'][3],
|
|
389
|
+
"cluster_number" : i['id']}) for i in text_blocks]
|
|
390
|
+
non_text_blocks = temp.pop('non_text_blocks')
|
|
391
|
+
temp['non_text_objects'] = [TextCluster.model_validate({"text" : i['text'],
|
|
392
|
+
"bb_y_min" : i['bounding_box'][0],
|
|
393
|
+
"bb_x_min" : i['bounding_box'][1],
|
|
394
|
+
"bb_y_max" : i['bounding_box'][2],
|
|
395
|
+
"bb_x_max" : i['bounding_box'][3],
|
|
396
|
+
"cluster_number" : i['id']}) for i in non_text_blocks]
|
|
397
|
+
return OCRClusterResponse.model_validate(temp)
|
|
398
|
+
from .utils.langchain import GeminiCostCallbackHandler
|
|
399
|
+
def refine_image_response(ok2, response_dict, outfile_name, image_file_path, model_names, cb: GeminiCostCallbackHandler):
|
|
400
|
+
|
|
401
|
+
# if allow_page_refine and (not preexisting):
|
|
402
|
+
if not response_dict:
|
|
403
|
+
with open(outfile_name, 'r') as f:
|
|
404
|
+
response_dict = json.load(f)
|
|
405
|
+
if response_dict.get('refined_pipeline_run'):
|
|
406
|
+
return False
|
|
407
|
+
if not retry_failed_refine and response_dict.get("refined_pipeline_failed_reason") is not None:
|
|
408
|
+
return False
|
|
409
|
+
cb.total_input_tokens = response_dict['usage_metadata']["input_tokens"]
|
|
410
|
+
cb.total_output_tokens = response_dict['usage_metadata']["output_tokens"]
|
|
411
|
+
cb.total_cost = response_dict['usage_metadata']["total_cost"]
|
|
412
|
+
cb.usage_history = response_dict['usage_metadata']["usage_history"]
|
|
413
|
+
i_model = [i for i, n in enumerate(model_names) if n.startswith('gemini-2.5')][0]
|
|
414
|
+
error_messages = []
|
|
415
|
+
refined = False
|
|
416
|
+
while not ok2:
|
|
417
|
+
refined = False
|
|
418
|
+
try:
|
|
419
|
+
|
|
420
|
+
model_name = model_names[i_model]
|
|
421
|
+
if "flash-lite" in model_name:
|
|
422
|
+
i_model += 1
|
|
423
|
+
if i_model >= min(len(model_names), 20):
|
|
424
|
+
logger.error(f"All LLM returned None as response, file name = {image_file_path}")
|
|
425
|
+
raise(ValueError(f"All LLM returned None as response, file name = {image_file_path}"))
|
|
426
|
+
continue
|
|
427
|
+
llm = _build_ocr_llm(model_name, callbacks=[cb])
|
|
428
|
+
|
|
429
|
+
refined = refine_table_ocr(response_dict, llm = llm, cb = cb, error_messages=error_messages)
|
|
430
|
+
ok2 = True
|
|
431
|
+
except Exception as e:
|
|
432
|
+
from .utils.log import safe_format_exception
|
|
433
|
+
e_prompt = safe_format_exception(e)
|
|
434
|
+
error_message = SystemMessage("post process error raised:\n"
|
|
435
|
+
f"{e_prompt[-10000:]}"
|
|
436
|
+
)
|
|
437
|
+
error_messages.append(error_message)
|
|
438
|
+
i_model += 1
|
|
439
|
+
refined = False
|
|
440
|
+
response_dict['refined_pipeline_failed_reason'] = "all model exhaused"
|
|
441
|
+
if i_model >= min(len(model_names), 20):
|
|
442
|
+
logger.error(f"All LLM returned None as response, file name = {image_file_path}")
|
|
443
|
+
ok2 = True # it is still ok even not refined
|
|
444
|
+
# raise(ValueError(f"All LLM returned None as response, file name = {image_file_path}"))
|
|
445
|
+
finally:
|
|
446
|
+
if refined:
|
|
447
|
+
response_dict['refined_version']['usage_metadata'] = cb.model_dump() # or usage_metadata
|
|
448
|
+
response_dict['refined_pipeline_run'] = True
|
|
449
|
+
with open(outfile_name, 'w') as f:
|
|
450
|
+
json.dump(response_dict, f)
|
|
451
|
+
|
|
452
|
+
print(response_dict)
|
|
453
|
+
return refined
|
|
454
|
+
def get_messages(image_file_path):
|
|
455
|
+
|
|
456
|
+
# Open the image in binary mode and read its content.
|
|
457
|
+
with open(image_file_path, "rb") as image_file:
|
|
458
|
+
image_bytes = image_file.read()
|
|
459
|
+
|
|
460
|
+
# Base64-encode the binary data.
|
|
461
|
+
encoded_bytes = base64.b64encode(image_bytes)
|
|
462
|
+
|
|
463
|
+
# Convert the encoded bytes to a UTF-8 string (optional, if you need a string representation)
|
|
464
|
+
encoded_str = encoded_bytes.decode('utf-8')
|
|
465
|
+
|
|
466
|
+
# Print the Base64-encoded string.
|
|
467
|
+
#print(encoded_str)
|
|
468
|
+
sys_message = SystemMessage("You are a helpful raw document AI that does OCR (Optical Character Reading) that focus on extracting raw document meta. "
|
|
469
|
+
"You must include spatial arrangement in the responded text. You must include bounding boxes locations x min, x max, y min and y max."
|
|
470
|
+
"The user provided image may contain partial document or some lost information on the edges. Try to recover as much information as possible. "
|
|
471
|
+
"The user provided image may contain paragraphs, figures or tables. Coerce your result to comply with the required json output format. "
|
|
472
|
+
# "If some error messages already found in other attempts, try to focus on parsing complicated tables. Remove watermarks. "
|
|
473
|
+
)
|
|
474
|
+
img_message = HumanMessage(
|
|
475
|
+
content=[
|
|
476
|
+
{"type": "text", "text": "find all text in the attached png file. "
|
|
477
|
+
},
|
|
478
|
+
{
|
|
479
|
+
"type": "image_url",
|
|
480
|
+
"image_url": {"url": f"data:image/png;base64,{encoded_str}"},
|
|
481
|
+
},
|
|
482
|
+
],
|
|
483
|
+
)
|
|
484
|
+
return sys_message, img_message
|
|
485
|
+
def ocr_single_image(gemini_key: str, page_file_name, file_name,
|
|
486
|
+
folder, # out folder
|
|
487
|
+
model_retry_priority_list : None | list[str],
|
|
488
|
+
exist_behavior: Literal["ok", "skip", "raise", 'rerun'] = 'skip'):
|
|
489
|
+
ok2 = False # stage 2 ok
|
|
490
|
+
outfile_name = os.path.join(folder,file_name, page_file_name.rsplit('.',1)[0] + '.json')
|
|
491
|
+
if os.path.exists(outfile_name):
|
|
492
|
+
if exist_behavior in ["ok", 'rerun']:
|
|
493
|
+
pass
|
|
494
|
+
elif exist_behavior == "skip":
|
|
495
|
+
pass
|
|
496
|
+
# return
|
|
497
|
+
else:
|
|
498
|
+
raise( PermissionError(f"output file {outfile_name} exists"))
|
|
499
|
+
|
|
500
|
+
assert gemini_key.startswith("AIza") # gcp keys
|
|
501
|
+
model_names = model_retry_priority_list or [# "gemini-2.0-flash", "gemini-2.0-flash-lite",
|
|
502
|
+
"gemini-3-flash-preview",
|
|
503
|
+
"gemini-2.5-flash",
|
|
504
|
+
"gemini-2.5-flash-lite", "gemini-2.5-pro",
|
|
505
|
+
# "gemini-1.5-pro",
|
|
506
|
+
"gemini-2.5-pro-preview-05-06", "gemini-2.5-pro-preview-03-25",
|
|
507
|
+
#"gemini-1.5-pro-latest", deprecated
|
|
508
|
+
"gemini-2.5-flash-preview-05-20"#, "gemini-1.5-flash"
|
|
509
|
+
# "gemini-2.5-flash-preview-04-17",
|
|
510
|
+
# "gemini-1.5-pro-001", "gemini-1.5-pro-002",
|
|
511
|
+
#"gemini-2.0-flash-thinking-exp-01-21",
|
|
512
|
+
#"gemini-1.5-flash",
|
|
513
|
+
]
|
|
514
|
+
draft_responses = {}
|
|
515
|
+
ok = False
|
|
516
|
+
i_model = 0
|
|
517
|
+
usage_metadata = []
|
|
518
|
+
from .utils.langchain import get_gemini_callback_cost
|
|
519
|
+
response_dict: dict = {}
|
|
520
|
+
image_file_path: str = os.path.join(folder, file_name, page_file_name)
|
|
521
|
+
with get_gemini_callback_cost() as cb:
|
|
522
|
+
|
|
523
|
+
if os.path.exists(outfile_name) and exist_behavior == 'rerun' or not os.path.exists(outfile_name):
|
|
524
|
+
sys_message, img_message = get_messages(image_file_path)
|
|
525
|
+
messages = [sys_message, img_message]
|
|
526
|
+
while not ok:
|
|
527
|
+
model_name = model_names[i_model]
|
|
528
|
+
try:
|
|
529
|
+
|
|
530
|
+
llm = _build_ocr_llm(model_name, callbacks=[cb])
|
|
531
|
+
response: OCRClusterResponse | None = get_first_round_response(draft_responses, llm, model_name, cb, messages, sys_message, img_message, usage_metadata)
|
|
532
|
+
_sp = validate_response_mutate_inplace(response, response_dict, image_file_path, model_name, page_file_name)
|
|
533
|
+
ok = True
|
|
534
|
+
# chain_new_raw = llm.with_structured_output(RawOCRResponse, include_raw = True)
|
|
535
|
+
except Exception as e:
|
|
536
|
+
from .utils.log import safe_format_exception
|
|
537
|
+
e_prompt = safe_format_exception(e)
|
|
538
|
+
error_message = SystemMessage("post process error raised:\n"
|
|
539
|
+
f"{e_prompt[-10000:]}"
|
|
540
|
+
)
|
|
541
|
+
messages.append(error_message)
|
|
542
|
+
i_model += 1
|
|
543
|
+
if i_model >= min(len(model_names), 20):
|
|
544
|
+
logger.error(f"All LLM returned None as response, file name = {image_file_path}")
|
|
545
|
+
final_resort(draft_responses, messages, page_file_name, model_name, image_file_path)
|
|
546
|
+
finally:
|
|
547
|
+
time.sleep(5)
|
|
548
|
+
assert response_dict, Exception("response_dict unbound")
|
|
549
|
+
|
|
550
|
+
response_dict['usage_metadata'] = cb.model_dump() # or usage_metadata
|
|
551
|
+
with open(outfile_name, 'w') as f:
|
|
552
|
+
json.dump(response_dict, f)
|
|
553
|
+
ok2 = False
|
|
554
|
+
print(response_dict)
|
|
555
|
+
allow_page_refine = False
|
|
556
|
+
if allow_page_refine:
|
|
557
|
+
refined = refine_image_response(ok2, response_dict, outfile_name, image_file_path, model_names, cb)
|
|
558
|
+
if refined:
|
|
559
|
+
time.sleep(5)
|
|
560
|
+
|
|
561
|
+
OCRRefineResponse: TypeAlias = OCRClusterResponse[DtoField]
|
|
562
|
+
def refine_table_ocr(response_dict, llm: BaseChatModel, cb, error_messages):
|
|
563
|
+
if response_dict.get('refined_version'):
|
|
564
|
+
return False
|
|
565
|
+
else:
|
|
566
|
+
response_dict['refined_version'] = None
|
|
567
|
+
sp= SplitPage.model_validate(response_dict)
|
|
568
|
+
if sp.refined_version:
|
|
569
|
+
return False
|
|
570
|
+
if not sp.contains_table:
|
|
571
|
+
return False
|
|
572
|
+
if len(sp.OCR_text_clusters) <5:
|
|
573
|
+
return False
|
|
574
|
+
system_prompt = SystemMessage("You are an OCR data organiser. Your job is to refine the user query that contains existing OCR raw text clusters to become meaningful. \n"
|
|
575
|
+
"For example, if a table cell or grid is broken down into multiple rows and was classified into multiple cells, \n"
|
|
576
|
+
"combine them all into a single meaningful text cluster. Keep metadata unchanged. Make sure no text is ever lost. \n"
|
|
577
|
+
"Include all punctionations, typos. Keep spelling errors. You must retain the text from original documents. \n"
|
|
578
|
+
"Beware that if an open quotation or open bracket is in one cluster and combine with another cluster, do keep those quotes or brackets.")
|
|
579
|
+
cluster_prompt = HumanMessage(f"{sp}")
|
|
580
|
+
|
|
581
|
+
# refine llm here\
|
|
582
|
+
messages = [system_prompt, cluster_prompt] + error_messages
|
|
583
|
+
max_attempt = 1 # nice to have, reorganise
|
|
584
|
+
for i in range(max_attempt):
|
|
585
|
+
try:
|
|
586
|
+
|
|
587
|
+
oc_refined_result: OCRRefineResponse
|
|
588
|
+
raw: str
|
|
589
|
+
parsing_error: Exception
|
|
590
|
+
temp: dict = cast(dict, llm.with_structured_output(schema = OCRRefineResponse, include_raw = True).invoke(messages, config={"callbacks": [cb]}))
|
|
591
|
+
(raw, oc_refined_result, parsing_error) = (temp['raw'], temp['parsed'], temp['parsing_error'])
|
|
592
|
+
if parsing_error:
|
|
593
|
+
raise parsing_error
|
|
594
|
+
text_before = [i.text for i in sp.OCR_text_clusters]
|
|
595
|
+
text_after = ' '.join([i.text for i in oc_refined_result.OCR_text_clusters])
|
|
596
|
+
from rapidfuzz import fuzz
|
|
597
|
+
def get_threshold(text_before):
|
|
598
|
+
if len(text_before) < 30:
|
|
599
|
+
threshold = 100
|
|
600
|
+
elif len(text_before) < 60:
|
|
601
|
+
threshold = 98
|
|
602
|
+
else:
|
|
603
|
+
threshold = 95
|
|
604
|
+
return threshold
|
|
605
|
+
is_preserved = [fuzz.partial_ratio(text_after, i) >= get_threshold(i) for i in text_before]
|
|
606
|
+
lost_text = [i for i, preserved in zip(sp.OCR_text_clusters, is_preserved) if not preserved]
|
|
607
|
+
ok = all(tf or (tc.text.strip() == "") for tf, tc in zip(is_preserved, sp.OCR_text_clusters))
|
|
608
|
+
assert ok, f"Some text or punctuations are lost through OCR text grouping, lost text = {str(lost_text)}"
|
|
609
|
+
response_dict['refined_version'] = oc_refined_result.model_dump()
|
|
610
|
+
break
|
|
611
|
+
except Exception as e:
|
|
612
|
+
if i < max_attempt-1:
|
|
613
|
+
messages.append(SystemMessage("error found: " + str(e)))
|
|
614
|
+
else:
|
|
615
|
+
raise
|
|
616
|
+
sp= SplitPage.model_validate(response_dict)
|
|
617
|
+
return True
|
|
618
|
+
|
|
619
|
+
|
|
620
|
+
def index_doc_group(doc_group_dumped):
|
|
621
|
+
doc_group_indexed = {(f,i['pdf_page_num']): i for f in doc_group_dumped['documents'] for i in doc_group_dumped['documents'][f]}
|
|
622
|
+
return doc_group_indexed
|
|
623
|
+
class Doc(BaseModel):
|
|
624
|
+
file_full_path: str = Field(description = "file full path")
|
|
625
|
+
pages: list[dict[str, Any]] = Field(description = "pages")
|
|
626
|
+
|
|
627
|
+
class DocumentGroup(BaseModel):
|
|
628
|
+
documents : dict[str, list|Doc] = Field(description = 'list of documents')
|
|
629
|
+
def to_doc_group_indexed(self):
|
|
630
|
+
doc_group = self.model_dump()
|
|
631
|
+
doc_group_indexed = index_doc_group(doc_group)
|
|
632
|
+
return doc_group_indexed
|
|
633
|
+
@staticmethod
|
|
634
|
+
def from_doc_folder(folder_path):
|
|
635
|
+
""" Assume the dir contains a list of folder with each folder with the filename
|
|
636
|
+
each subfolder contains a list of pages
|
|
637
|
+
|
|
638
|
+
Args:
|
|
639
|
+
folder_path (_type_): _description_
|
|
640
|
+
"""
|
|
641
|
+
dirs = os.listdir(folder_path)
|
|
642
|
+
doc_group = {}
|
|
643
|
+
for d in dirs:
|
|
644
|
+
doc = regen_doc(os.path.join(folder_path, d))
|
|
645
|
+
doc_group[d] = doc
|
|
646
|
+
return DocumentGroup(**{"documents": doc_group})
|
|
647
|
+
def regen_doc_group(folder_path, use_raw = False):
|
|
648
|
+
""" Assume the dir contains a list of folder with each folder with the filename
|
|
649
|
+
each subfolder contains a list of pages
|
|
650
|
+
|
|
651
|
+
Args:
|
|
652
|
+
folder_path (_type_): _description_
|
|
653
|
+
"""
|
|
654
|
+
dirs = os.listdir(folder_path)
|
|
655
|
+
doc_group = {}
|
|
656
|
+
for d in dirs:
|
|
657
|
+
doc = regen_doc(os.path.join(folder_path, d), use_raw = use_raw)
|
|
658
|
+
doc_group[d] = doc
|
|
659
|
+
return DocumentGroup(**{"documents": doc_group})
|
|
660
|
+
|
|
661
|
+
try:
|
|
662
|
+
from .utils.bounded_threadpool_executor import BoundedExecutor
|
|
663
|
+
except ImportError: # pragma: no cover
|
|
664
|
+
from .utils.bounded_threadpool_executor import BoundedExecutor
|
|
665
|
+
|
|
666
|
+
|
|
667
|
+
def get_legacy_loader_like(folder: str, allowed_relative_paths: str | Any):
|
|
668
|
+
|
|
669
|
+
|
|
670
|
+
# useful for old flat entry only
|
|
671
|
+
directories = [entry for entry in os.listdir(folder) if os.path.isdir(os.path.join(folder, entry))]
|
|
672
|
+
if allowed_relative_paths is None:
|
|
673
|
+
if os.path.exists("allowed_file_name_OCR.txt"):
|
|
674
|
+
import json
|
|
675
|
+
with open("allowed_file_name_OCR.txt", 'r') as f:
|
|
676
|
+
allowed_files = json.load(f)
|
|
677
|
+
else:
|
|
678
|
+
allowed_files = directories
|
|
679
|
+
else:
|
|
680
|
+
# fix for non flat later cases tree structure
|
|
681
|
+
allowed_files = allowed_relative_paths
|
|
682
|
+
def local_loader():
|
|
683
|
+
for root, dirs, files in os.walk(folder):
|
|
684
|
+
import pathlib
|
|
685
|
+
if root==folder:
|
|
686
|
+
continue
|
|
687
|
+
(pdf_folder, pdf_fname) = os.path.split(root)
|
|
688
|
+
if str(pathlib.Path(root).relative_to(folder)) in allowed_files:
|
|
689
|
+
pass
|
|
690
|
+
else:
|
|
691
|
+
continue
|
|
692
|
+
if dirs != []:
|
|
693
|
+
continue
|
|
694
|
+
for f in files:
|
|
695
|
+
full_path = os.path.join(root, f)
|
|
696
|
+
page_file_name = f
|
|
697
|
+
if page_file_name.endswith('.png') and not os.path.exists(os.path.join(folder, pdf_fname, page_file_name.rsplit('.',1)[0] + ".json")):
|
|
698
|
+
pass
|
|
699
|
+
else:
|
|
700
|
+
continue
|
|
701
|
+
yield page_file_name
|
|
702
|
+
return local_loader()
|
|
703
|
+
|
|
704
|
+
def batch_gemini_ocr_image(gemini_key, folder = "split_pages", exist_behavior: Literal["ok","skip","raise", 'rerun'] = 'skip',
|
|
705
|
+
bounded_executor: Optional[BoundedExecutor] = None,
|
|
706
|
+
allowed_relative_paths = None,
|
|
707
|
+
loader : RawFileLoader | None= None,
|
|
708
|
+
ocr_callback : Callable| None= None):
|
|
709
|
+
# page_file_name = "page_1.png"
|
|
710
|
+
# file_name = "EXL-00-HI-MSA01-2017.PDF"
|
|
711
|
+
|
|
712
|
+
if loader:
|
|
713
|
+
pdf_folder = loader.compare_root
|
|
714
|
+
legacy_mode = False
|
|
715
|
+
loader2 = loader
|
|
716
|
+
else:
|
|
717
|
+
legacy_mode = True
|
|
718
|
+
|
|
719
|
+
loader2: Iterable = get_legacy_loader_like(folder, allowed_relative_paths)
|
|
720
|
+
for f in loader2:
|
|
721
|
+
import pathlib
|
|
722
|
+
#print(f'start inspecting processing {f}')
|
|
723
|
+
logging.info(f'start inspecting processing {f}')
|
|
724
|
+
if legacy_mode:
|
|
725
|
+
pdf_folder = folder
|
|
726
|
+
page_file_name = pathlib.Path(f).parts[-1]
|
|
727
|
+
pdf_fname = pathlib.Path(f).parts[-2]
|
|
728
|
+
else:
|
|
729
|
+
|
|
730
|
+
if hasattr(loader, "compare_root") and type(loader) is RawFileLoader:
|
|
731
|
+
pdf_folder = str(pathlib.Path(os.path.join(loader.compare_root, f)).parent.parent)
|
|
732
|
+
page_file_name = pathlib.Path(f).parts[-1]
|
|
733
|
+
pdf_fname = pathlib.Path(f).parts[-2]
|
|
734
|
+
else:
|
|
735
|
+
raise Exception("Unreachable")
|
|
736
|
+
if bounded_executor is None:
|
|
737
|
+
(ocr_callback or ocr_single_image)(gemini_key, page_file_name,
|
|
738
|
+
file_name=pdf_fname,
|
|
739
|
+
folder=pdf_folder,
|
|
740
|
+
exist_behavior=exist_behavior,
|
|
741
|
+
model_retry_priority_list=None)
|
|
742
|
+
else:
|
|
743
|
+
bounded_executor.submit(ocr_callback or ocr_single_image,
|
|
744
|
+
gemini_key,
|
|
745
|
+
page_file_name,
|
|
746
|
+
file_name=pdf_fname,
|
|
747
|
+
folder=pdf_folder,
|
|
748
|
+
exist_behavior=exist_behavior)
|
|
749
|
+
time.sleep(2)
|
|
750
|
+
if bounded_executor:
|
|
751
|
+
bounded_executor.wait_for_all()
|
|
752
|
+
bounded_executor.shutdown()
|