Spaces:
Running on Zero
Running on Zero
Download molparser/models/molparser.py from hugging-apps/molparser-mobile-v2: direct link, hf CLI and curl.
- Browser
- Download file 16.4 kB
-
https://huggingface.co/spaces/hugging-apps/molparser-mobile-v2/resolve/main/molparser/models/molparser.py
- Command line
-
hf download hf://spaces/hugging-apps/molparser-mobile-v2/molparser/models/molparser.py
-
curl -L -o molparser.py https://huggingface.co/spaces/hugging-apps/molparser-mobile-v2/resolve/main/molparser/models/molparser.py
16.4 kB
| from __future__ import annotations | |
| import os | |
| from contextlib import contextmanager | |
| from dataclasses import asdict, dataclass, fields, replace | |
| from importlib.resources import files | |
| from pathlib import Path | |
| from typing import Any, Optional | |
| from .io import load_image, normalize_inputs, render_pdf | |
| from .runtime import resolve_model | |
| from ..utils import postprocess_caption | |
| DEFAULT_CONFIG = "config.yaml" | |
| class MolParserConfig: | |
| moldet_hf_repo: str = "UniParser/MolDetv2" | |
| moldet_modelscope_repo: str = "UniParser/MolDetv2" | |
| moldet_image_model_path: str = "" | |
| moldet_pdf_model_path: str = "" | |
| moldet_image_modelname: str = "moldet_v2_yolo11n_640_general.pt" | |
| moldet_pdf_modelname: str = "moldet_v2_yolo11n_960_doc.pt" | |
| molparser_hf_repo: str = "UniParser/MolParser-Mobile" | |
| molparser_modelscope_repo: str = "UniParser/MolParser-Mobile" | |
| molparser_model_path: str = "" | |
| cache_dir: str = "" | |
| device: str = "auto" | |
| detector_conf: float = 0.5 | |
| image_imgsz: int = 640 | |
| pdf_imgsz: int = 960 | |
| moldet_batch_size: int = 8 | |
| pdf_dpi: int = 200 | |
| max_length: int = 256 | |
| molparser_batch_size: int = 32 | |
| expand_px: int = 2 | |
| padding_px: int = 0 | |
| hf_token: str | bool | None = None | |
| def from_default(cls, **overrides) -> "MolParserConfig": | |
| return cls.from_yaml(files("molparser.models").joinpath(DEFAULT_CONFIG), **overrides) | |
| def from_yaml(cls, path: str | Path, **overrides) -> "MolParserConfig": | |
| try: | |
| import yaml | |
| except ImportError as exc: | |
| raise ImportError("YAML config loading requires PyYAML.") from exc | |
| with Path(path).open("r", encoding="utf-8") as f: | |
| data = yaml.safe_load(f) or {} | |
| return cls.from_mapping(data, **overrides) | |
| def from_mapping(cls, data: dict[str, Any], **overrides) -> "MolParserConfig": | |
| flat = { | |
| "moldet_hf_repo": data.get("moldet", {}).get("hf_repo", cls.moldet_hf_repo), | |
| "moldet_modelscope_repo": data.get("moldet", {}).get("modelscope_repo", cls.moldet_modelscope_repo), | |
| "moldet_image_model_path": data.get("moldet", {}).get("image_model_path", ""), | |
| "moldet_pdf_model_path": data.get("moldet", {}).get("pdf_model_path", ""), | |
| "moldet_image_modelname": data.get("moldet", {}).get("image_modelname", cls.moldet_image_modelname), | |
| "moldet_pdf_modelname": data.get("moldet", {}).get("pdf_modelname", cls.moldet_pdf_modelname), | |
| "detector_conf": data.get("moldet", {}).get("confidence", cls.detector_conf), | |
| "image_imgsz": data.get("moldet", {}).get("image_imgsz", cls.image_imgsz), | |
| "pdf_imgsz": data.get("moldet", {}).get("pdf_imgsz", cls.pdf_imgsz), | |
| "moldet_batch_size": data.get("moldet", {}).get("batch_size", cls.moldet_batch_size), | |
| "expand_px": data.get("moldet", {}).get("expand_px", cls.expand_px), | |
| "molparser_hf_repo": data.get("molparser", {}).get("hf_repo", cls.molparser_hf_repo), | |
| "molparser_modelscope_repo": data.get("molparser", {}).get("modelscope_repo", ""), | |
| "molparser_model_path": data.get("molparser", {}).get("model_path", ""), | |
| "max_length": data.get("molparser", {}).get("max_length", cls.max_length), | |
| "molparser_batch_size": data.get("molparser", {}).get("batch_size", cls.molparser_batch_size), | |
| "padding_px": data.get("molparser", {}).get("padding_px", cls.padding_px), | |
| "pdf_dpi": data.get("pdf", {}).get("dpi", cls.pdf_dpi), | |
| "device": data.get("runtime", {}).get("device", cls.device), | |
| "cache_dir": data.get("runtime", {}).get("cache_dir", ""), | |
| "hf_token": data.get("runtime", {}).get("hf_token"), | |
| } | |
| flat.update(_normalize_overrides(overrides)) | |
| allowed = {field.name for field in fields(cls)} | |
| return cls(**{key: value for key, value in flat.items() if key in allowed}) | |
| class MolParserResult: | |
| source: str | |
| input_index: int | |
| page_index: Optional[int] | |
| bbox: Optional[tuple[float, float, float, float]] | |
| confidence: Optional[float] | |
| raw_caption: str | |
| caption: str | |
| smi: str | |
| esmi: str | |
| cxsmiles: str | |
| markush: bool | |
| sru: bool | |
| groups: Any | |
| def to_dict(self) -> dict[str, Any]: | |
| return asdict(self) | |
| class MolParser: | |
| def __init__(self, config: MolParserConfig | str | Path | None = None, **overrides): | |
| if config is None: | |
| self.config = MolParserConfig.from_default(**overrides) | |
| elif isinstance(config, (str, Path)): | |
| self.config = MolParserConfig.from_yaml(config, **overrides) | |
| elif overrides: | |
| self.config = replace(config, **_normalize_overrides(overrides)) | |
| else: | |
| self.config = config | |
| self._molparser = None | |
| self._image_detector = None | |
| self._pdf_detector = None | |
| def parse( | |
| self, | |
| inputs, | |
| *, | |
| rec_only: bool = False, | |
| pages=None, | |
| expand_px: int | None = None, | |
| padding_px: int | None = None, | |
| ) -> list[MolParserResult]: | |
| with self._temporary_margins(expand_px=expand_px, padding_px=padding_px): | |
| image_records: list[tuple[Any, str, int]] = [] | |
| pdf_records: list[tuple[Path, str, int]] = [] | |
| for input_index, item in enumerate(normalize_inputs(inputs)): | |
| if item.kind == "pdf": | |
| if item.path is None: | |
| raise ValueError("PDF input requires a file path.") | |
| pdf_records.append((item.path, item.source, input_index)) | |
| else: | |
| image_records.append((load_image(item), item.source, input_index)) | |
| results: list[MolParserResult] = [] | |
| if image_records: | |
| results.extend(self._parse_image_batch(image_records, rec_only=rec_only)) | |
| if pdf_records: | |
| results.extend(self._parse_pdf_batch(pdf_records, pages=pages)) | |
| return sorted(results, key=lambda result: result.input_index) | |
| def parse_image( | |
| self, | |
| image, | |
| *, | |
| rec_only: bool = True, | |
| expand_px: int | None = None, | |
| padding_px: int | None = None, | |
| ) -> list[MolParserResult]: | |
| with self._temporary_margins(expand_px=expand_px, padding_px=padding_px): | |
| item = normalize_inputs(image)[0] | |
| return self._parse_image(load_image(item), item.source, 0, rec_only=rec_only) | |
| def parse_pdf( | |
| self, | |
| pdf, | |
| *, | |
| pages=None, | |
| expand_px: int | None = None, | |
| padding_px: int | None = None, | |
| ) -> list[MolParserResult]: | |
| with self._temporary_margins(expand_px=expand_px, padding_px=padding_px): | |
| item = normalize_inputs(pdf)[0] | |
| if item.kind != "pdf" or item.path is None: | |
| raise ValueError("parse_pdf expects a PDF path or URL.") | |
| return self._parse_pdf_batch([(item.path, item.source, 0)], pages=pages) | |
| def _token(self): | |
| return self.config.hf_token if self.config.hf_token is not None else os.environ.get("HF_TOKEN") | |
| def _cache_dir(self): | |
| return self.config.cache_dir or None | |
| def _temporary_margins(self, *, expand_px: int | None, padding_px: int | None): | |
| old_expand = self.config.expand_px | |
| old_padding = self.config.padding_px | |
| if expand_px is not None: | |
| self.config.expand_px = int(expand_px) | |
| if padding_px is not None: | |
| self.config.padding_px = int(padding_px) | |
| try: | |
| yield | |
| finally: | |
| self.config.expand_px = old_expand | |
| self.config.padding_px = old_padding | |
| def _recognize(self, images: list[Any]) -> list[str]: | |
| if self._molparser is None: | |
| from .runtime import MolParserRecognizer | |
| path = resolve_model( | |
| local_path=self.config.molparser_model_path, | |
| hf_model_id=self.config.molparser_hf_repo, | |
| modelscope_model_id=self.config.molparser_modelscope_repo, | |
| cache_dir=self._cache_dir(), | |
| token=self._token(), | |
| ) | |
| self._molparser = MolParserRecognizer( | |
| str(path), | |
| device=self.config.device, | |
| token=self._token(), | |
| max_length=self.config.max_length, | |
| ) | |
| captions: list[str] = [] | |
| for batch in _batched(images, self.config.molparser_batch_size): | |
| captions.extend(self._molparser.recognize(batch)) | |
| return captions | |
| def _detector(self, kind: str): | |
| from .runtime import MolDetDetector | |
| if kind == "pdf": | |
| if self._pdf_detector is None: | |
| path = self._resolve_detector(self.config.moldet_pdf_model_path, self.config.moldet_pdf_modelname) | |
| self._pdf_detector = MolDetDetector( | |
| str(path), device=self.config.device, imgsz=self.config.pdf_imgsz, conf=self.config.detector_conf | |
| ) | |
| return self._pdf_detector | |
| if self._image_detector is None: | |
| path = self._resolve_detector(self.config.moldet_image_model_path, self.config.moldet_image_modelname) | |
| self._image_detector = MolDetDetector( | |
| str(path), device=self.config.device, imgsz=self.config.image_imgsz, conf=self.config.detector_conf | |
| ) | |
| return self._image_detector | |
| def _resolve_detector(self, local_path: str, filename: str) -> Path: | |
| return resolve_model( | |
| local_path=local_path, | |
| hf_model_id=self.config.moldet_hf_repo, | |
| modelscope_model_id=self.config.moldet_modelscope_repo, | |
| filename=filename, | |
| cache_dir=self._cache_dir(), | |
| token=self._token(), | |
| ) | |
| def _parse_image(self, image, source: str, input_index: int, *, rec_only: bool) -> list[MolParserResult]: | |
| return self._parse_image_batch([(image, source, input_index)], rec_only=rec_only) | |
| def _parse_image_batch(self, image_records: list[tuple[Any, str, int]], *, rec_only: bool) -> list[MolParserResult]: | |
| images = [record[0] for record in image_records] | |
| sources = [record[1] for record in image_records] | |
| input_indexes = [record[2] for record in image_records] | |
| if rec_only: | |
| padded_images = [self._pad_image(image) for image in images] | |
| return self._results_many(padded_images, sources, input_indexes, [None] * len(images), [None] * len(images), [None] * len(images)) | |
| detections_by_image = self._detect("image", images) | |
| crops: list[Any] = [] | |
| crop_sources: list[str] = [] | |
| crop_input_indexes: list[int] = [] | |
| bboxes: list[tuple[float, float, float, float] | None] = [] | |
| confidences: list[float | None] = [] | |
| for image, source, input_index, detections in zip(images, sources, input_indexes, detections_by_image): | |
| for detection in detections: | |
| crops.append(self._crop(image, detection)) | |
| crop_sources.append(source) | |
| crop_input_indexes.append(input_index) | |
| bboxes.append(detection.bbox) | |
| confidences.append(detection.confidence) | |
| return self._results_many( | |
| crops, | |
| crop_sources, | |
| crop_input_indexes, | |
| [None] * len(crops), | |
| bboxes, | |
| confidences, | |
| ) | |
| def _parse_pdf(self, path: Path, source: str, input_index: int, *, pages=None) -> list[MolParserResult]: | |
| return self._parse_pdf_batch([(path, source, input_index)], pages=pages) | |
| def _parse_pdf_batch(self, pdf_records: list[tuple[Path, str, int]], *, pages=None) -> list[MolParserResult]: | |
| page_records = [] | |
| for path, source, input_index in pdf_records: | |
| for page in render_pdf(path, dpi=self.config.pdf_dpi, pages=pages): | |
| page_records.append((page.image, source, input_index, page.page_index)) | |
| if not page_records: | |
| return [] | |
| page_images = [record[0] for record in page_records] | |
| detections_by_page = self._detect("pdf", page_images) | |
| crops: list[Any] = [] | |
| crop_sources: list[str] = [] | |
| crop_input_indexes: list[int] = [] | |
| page_indexes: list[int | None] = [] | |
| bboxes: list[tuple[float, float, float, float] | None] = [] | |
| confidences: list[float | None] = [] | |
| for (image, source, input_index, page_index), detections in zip(page_records, detections_by_page): | |
| for detection in detections: | |
| crops.append(self._crop(image, detection)) | |
| crop_sources.append(source) | |
| crop_input_indexes.append(input_index) | |
| page_indexes.append(page_index) | |
| bboxes.append(detection.bbox) | |
| confidences.append(detection.confidence) | |
| return self._results_many(crops, crop_sources, crop_input_indexes, page_indexes, bboxes, confidences) | |
| def _crop(self, image, detection): | |
| from .runtime import crop_detection | |
| return crop_detection(image, detection, expand_px=self.config.expand_px, pad_px=self.config.padding_px) | |
| def _pad_image(self, image): | |
| from .runtime import pad_image | |
| return pad_image(image, pad_px=self.config.padding_px) | |
| def _detect(self, kind: str, images: list[Any]): | |
| detector = self._detector(kind) | |
| detections: list[Any] = [] | |
| for batch in _batched(images, self.config.moldet_batch_size): | |
| detections.extend(detector.detect(batch)) | |
| return detections | |
| def _results( | |
| self, | |
| images: list[Any], | |
| source: str, | |
| input_index: int, | |
| page_indexes: list[int | None], | |
| bboxes: list[tuple[float, float, float, float] | None], | |
| confidences: list[float | None], | |
| ) -> list[MolParserResult]: | |
| return self._results_many( | |
| images, | |
| [source] * len(images), | |
| [input_index] * len(images), | |
| page_indexes, | |
| bboxes, | |
| confidences, | |
| ) | |
| def _results_many( | |
| self, | |
| images: list[Any], | |
| sources: list[str], | |
| input_indexes: list[int], | |
| page_indexes: list[int | None], | |
| bboxes: list[tuple[float, float, float, float] | None], | |
| confidences: list[float | None], | |
| ) -> list[MolParserResult]: | |
| if not images: | |
| return [] | |
| raws = self._recognize(images) | |
| results: list[MolParserResult] = [] | |
| for raw, source, input_index, page_index, bbox, confidence in zip(raws, sources, input_indexes, page_indexes, bboxes, confidences): | |
| post = postprocess_caption(raw) | |
| results.append( | |
| MolParserResult( | |
| source=source, | |
| input_index=input_index, | |
| page_index=page_index, | |
| bbox=bbox, | |
| confidence=confidence, | |
| raw_caption=str(raw), | |
| caption=str(post.get("caption", raw)), | |
| smi=str(post.get("smi", "")), | |
| esmi=str(post.get("esmi", "")), | |
| cxsmiles=str(post.get("cxsmiles", "")), | |
| markush=bool(post.get("markush", False)), | |
| sru=bool(post.get("sru", False)), | |
| groups=post.get("groups", ""), | |
| ) | |
| ) | |
| return results | |
| def _normalize_overrides(overrides: dict[str, Any]) -> dict[str, Any]: | |
| aliases = { | |
| "moldet_hf_model_id": "moldet_hf_repo", | |
| "moldet_modelscope_model_id": "moldet_modelscope_repo", | |
| "recognizer_hf_model_id": "molparser_hf_repo", | |
| "recognizer_modelscope_model_id": "molparser_modelscope_repo", | |
| "recognizer_model_path": "molparser_model_path", | |
| "det_batch_size": "moldet_batch_size", | |
| "ocsr_batch_size": "molparser_batch_size", | |
| } | |
| return {aliases.get(key, key): value for key, value in overrides.items()} | |
| def _batched(items: list[Any], batch_size: int): | |
| size = max(1, int(batch_size or 1)) | |
| for start in range(0, len(items), size): | |
| yield items[start : start + size] | |
| __all__ = ["MolParser", "MolParserConfig", "MolParserResult"] | |