Cesium2 / src /document.py
MORPH-AI
feat: dynamic MoE expansion, multi-head CoT, plugin architecture, improved MoD
82f262a
Raw
History Blame Contribute Delete
5.03 kB
"""
DocumentModule - PDF/DOCX/OCR with layout-aware parsing for MORPH-AI v6.
"""
import re
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from architecture import MorphConfig
@dataclass
class DocumentFacts:
text: str = ""
pages: int = 0
format: str = ""
tables: List[List[List[str]]] = field(default_factory=list)
metadata: Dict[str, str] = field(default_factory=dict)
embedding: Optional[torch.Tensor] = None
def to_text(self) -> str:
parts = [f"document {self.format} {self.pages}p"]
if self.text:
parts.append(f"text: {self.text[:500]}")
if self.tables:
parts.append(f"tables: {len(self.tables)}")
return " | ".join(parts)
def to_dict(self) -> Dict[str, Any]:
return {
"text": self.text,
"pages": self.pages,
"format": self.format,
"tables": self.tables,
"metadata": self.metadata,
}
class DocumentModule(nn.Module):
"""PDF/DOCX/OCR with layout-aware parsing for document understanding."""
def __init__(self, config: MorphConfig, hidden_dim: int):
super().__init__()
self.max_pages = config.doc_max_pages
self.page_proj = nn.Linear(hidden_dim, config.doc_hidden)
self.layout_encoder = nn.Sequential(
nn.Linear(config.doc_hidden + 4, config.doc_hidden),
nn.GELU(),
nn.Linear(config.doc_hidden, hidden_dim),
)
nn.init.zeros_(self.layout_encoder[-1].weight)
nn.init.zeros_(self.layout_encoder[-1].bias)
def forward(self, hidden: torch.Tensor, layout_info: Optional[torch.Tensor] = None) -> torch.Tensor:
B, T, H = hidden.shape
page_emb = self.page_proj(hidden)
if layout_info is not None:
layout = layout_info.to(hidden.dtype)
page_emb = self.layout_encoder(torch.cat([page_emb, layout], dim=-1))
return hidden + page_emb
def extract_text(self, source) -> str:
"""Extract text from PDF/DOCX/image with OCR fallback."""
try:
if hasattr(source, 'endswith'):
if source.endswith('.pdf'):
return self._extract_pdf(source)
elif source.endswith('.docx'):
return self._extract_docx(source)
elif source.endswith(('.png', '.jpg', '.jpeg', '.bmp', '.tiff')):
return self._extract_image_ocr(source)
return self._extract_image_ocr(source)
except Exception as e:
return f"[document extraction error: {e}]"
def _extract_pdf(self, path: str) -> str:
try:
import fitz
doc = fitz.open(path)
pages = []
for i in range(min(len(doc), self.max_pages)):
page = doc[i]
text = page.get_text()
tables = page.find_tables()
if tables.tables:
for table in tables.tables:
pages.append(f"[TABLE]\n{table.to_pandas().to_string()}")
pages.append(text)
return "\n\n".join(pages)
except ImportError:
return "[PDF extraction requires PyMuPDF: pip install pymupdf]"
def _extract_docx(self, path: str) -> str:
try:
import docx2txt
return docx2txt.process(path)
except ImportError:
return "[DOCX extraction requires docx2txt: pip install docx2txt]"
def _extract_image_ocr(self, source) -> str:
try:
import pytesseract
from PIL import Image
img = Image.open(source)
text = pytesseract.image_to_string(img)
data = pytesseract.image_to_data(img, output_type=pytesseract.Output.DICT)
lines = []
for i, word in enumerate(data['text']):
if word.strip():
lines.append(word)
return text + "\n\n[LAYOUT]\n" + " ".join(lines)
except ImportError:
return "[OCR requires pytesseract + Pillow: pip install pytesseract pillow]"
except Exception as e:
return f"[OCR error: {e}]"
def analyze(self, source) -> DocumentFacts:
"""Full document analysis returning structured facts."""
facts = DocumentFacts()
try:
text = self.extract_text(source)
facts.text = text
facts.pages = len(text.split('\n\n'))
if hasattr(source, 'endswith'):
if source.endswith('.pdf'):
facts.format = 'PDF'
elif source.endswith('.docx'):
facts.format = 'DOCX'
else:
facts.format = 'IMAGE'
else:
facts.format = 'UNKNOWN'
except Exception as e:
facts.text = f"[analysis error: {e}]"
return facts