Spaces:
Running on Zero
Running on Zero
Download satquery_agent/model_artifacts.py from AnirudhShashikumar/SatQuery-AI: direct link, hf CLI and curl.
- Browser
- Download file 11.7 kB
-
https://huggingface.co/spaces/AnirudhShashikumar/SatQuery-AI/resolve/main/satquery_agent/model_artifacts.py
- Command line
-
hf download hf://spaces/AnirudhShashikumar/SatQuery-AI/satquery_agent/model_artifacts.py
-
curl -L -o model_artifacts.py https://huggingface.co/spaces/AnirudhShashikumar/SatQuery-AI/resolve/main/satquery_agent/model_artifacts.py
11.7 kB
| """Local-first resolution for the private CodeBlueX model artifacts.""" | |
| from __future__ import annotations | |
| import hashlib | |
| import os | |
| import shutil | |
| import threading | |
| from contextlib import contextmanager | |
| from dataclasses import dataclass | |
| from pathlib import Path, PurePosixPath | |
| from types import MappingProxyType | |
| from typing import Any, Callable, Iterator, Mapping, Optional | |
| PRIVATE_MODEL_REPOSITORY = "AnirudhShashikumar/SatQuery_AI_Models" | |
| SATQUERY_VISION_ENCODER_V1 = "satquery_vision_encoder_v1" | |
| RSVQA_SPECIALIST_V1 = "rsvqa_specialist_v1" | |
| GROUNDING_SPECIALIST_V1_1 = "grounding_specialist_v1_1" | |
| PIX2PIX_SAR_TRANSLATION = "pix2pix_sar_translation" | |
| SARFUSIONFORMER_SAR_TRANSLATION = "sarfusionformer_sar_translation" | |
| SARFUSIONFORMER_COLOR_CORRECTOR = "sarfusionformer_color_corrector" | |
| class ModelArtifact: | |
| key: str | |
| name: str | |
| repository_path: str | |
| sha256: str | |
| CODEBLUEX_ARTIFACTS: Mapping[str, ModelArtifact] = MappingProxyType( | |
| { | |
| SATQUERY_VISION_ENCODER_V1: ModelArtifact( | |
| key=SATQUERY_VISION_ENCODER_V1, | |
| name="SatQuery Vision Encoder v1 adapter", | |
| repository_path="satquery_vision_encoder_v1/satquery_vision_encoder_v1_adapter.pt", | |
| sha256="a99c0bf0fb44044988ef1698483888c8a2e3a047d2d2d56478837575cf7626ea", | |
| ), | |
| RSVQA_SPECIALIST_V1: ModelArtifact( | |
| key=RSVQA_SPECIALIST_V1, | |
| name="RSVQA Specialist v1 head", | |
| repository_path="rsvqa_specialist_v1/rsvqa_specialist_v1_head.pt", | |
| sha256="71c0ab56ee650813bd495e8a3bc777353b6907a097af860e417f60523efe56ad", | |
| ), | |
| GROUNDING_SPECIALIST_V1_1: ModelArtifact( | |
| key=GROUNDING_SPECIALIST_V1_1, | |
| name="Grounding Specialist v1.1 head", | |
| repository_path="grounding_specialist_v1_1/grounding_specialist_v1_1_head.pt", | |
| sha256="5e8db30becadb1d063fc0154614ee2a2a4b7a2c2923db3fce4ff036c65007342", | |
| ), | |
| PIX2PIX_SAR_TRANSLATION: ModelArtifact( | |
| key=PIX2PIX_SAR_TRANSLATION, | |
| name="Pix2Pix SAR reconstruction checkpoint", | |
| repository_path="sar_translation/pix2pix/pix2pix_gen_180.pth", | |
| sha256="5bdc6bf9b29986860d6537265280bb47b74f85ebe1222256ddc59acce3a743b3", | |
| ), | |
| SARFUSIONFORMER_SAR_TRANSLATION: ModelArtifact( | |
| key=SARFUSIONFORMER_SAR_TRANSLATION, | |
| name="SARFusionFormer SAR reconstruction checkpoint", | |
| repository_path="sar_translation/sarfusionformer/sarfusionformer_256_decoder_best.pt", | |
| sha256="6c8ac2d482b66877a910a9c95e6398f9e0586f13f786d0243c736a2d10b30548", | |
| ), | |
| SARFUSIONFORMER_COLOR_CORRECTOR: ModelArtifact( | |
| key=SARFUSIONFORMER_COLOR_CORRECTOR, | |
| name="SARFusionFormer color-corrector checkpoint", | |
| repository_path="sar_translation/sarfusionformer/color_corrector_256_best.pt", | |
| sha256="15298976fbffe93c16d81716c40fe3c7a09499e63e649edd4de9a1fd70ec675d", | |
| ), | |
| } | |
| ) | |
| class ModelArtifactError(RuntimeError): | |
| """Base class for credential-safe artifact resolution failures.""" | |
| class UnknownModelArtifactError(ModelArtifactError): | |
| """Raised when a caller requests an artifact outside the CodeBlueX allowlist.""" | |
| class ModelArtifactChecksumError(ModelArtifactError): | |
| """Raised when an existing or downloaded artifact has the wrong identity.""" | |
| class ModelArtifactAuthenticationError(ModelArtifactError): | |
| """Raised when a private download is needed but no environment token exists.""" | |
| class ModelArtifactMissingError(ModelArtifactError): | |
| """Raised when an explicit local override is absent and cannot be replaced.""" | |
| class ModelArtifactDownloadError(ModelArtifactError): | |
| """Raised when the private repository cannot provide the requested artifact.""" | |
| class ModelArtifactFilesystemError(ModelArtifactError): | |
| """Raised when a verified artifact cannot be published to its local path.""" | |
| DownloadFunction = Callable[..., Any] | |
| def sha256_file(path: Path) -> str: | |
| digest = hashlib.sha256() | |
| try: | |
| with path.open("rb") as stream: | |
| for chunk in iter(lambda: stream.read(1024 * 1024), b""): | |
| digest.update(chunk) | |
| except OSError as error: | |
| raise ModelArtifactFilesystemError(f"Could not read model artifact {path.name}.") from error | |
| return digest.hexdigest() | |
| def _exclusive_file_lock(path: Path) -> Iterator[None]: | |
| """Serialize publication across workers while keeping local reads lock-free.""" | |
| try: | |
| import fcntl | |
| with path.open("a+b") as lock_stream: | |
| fcntl.flock(lock_stream.fileno(), fcntl.LOCK_EX) | |
| try: | |
| yield | |
| finally: | |
| fcntl.flock(lock_stream.fileno(), fcntl.LOCK_UN) | |
| except ModelArtifactError: | |
| raise | |
| except OSError as error: | |
| raise ModelArtifactFilesystemError("Could not lock the local model artifact directory.") from error | |
| class ModelArtifactResolver: | |
| """Resolve pinned artifacts locally, downloading only an absent allowlisted file.""" | |
| def __init__( | |
| self, | |
| *, | |
| repository_id: str = PRIVATE_MODEL_REPOSITORY, | |
| artifacts: Mapping[str, ModelArtifact] = CODEBLUEX_ARTIFACTS, | |
| environment: Optional[Mapping[str, str]] = None, | |
| downloader: Optional[DownloadFunction] = None, | |
| models_root: Optional[Path] = None, | |
| ) -> None: | |
| self.repository_id = repository_id | |
| self.artifacts = dict(artifacts) | |
| self.environment = environment if environment is not None else os.environ | |
| self.downloader = downloader | |
| self.models_root = (models_root or Path(__file__).resolve().parents[1] / "models").resolve() | |
| def _artifact(self, key: str) -> ModelArtifact: | |
| try: | |
| artifact = self.artifacts[key] | |
| except KeyError: | |
| raise UnknownModelArtifactError( | |
| f"{key!r} is not an allowlisted CodeBlueX model artifact." | |
| ) from None | |
| repository_path = PurePosixPath(artifact.repository_path) | |
| if repository_path.is_absolute() or ".." in repository_path.parts: | |
| raise UnknownModelArtifactError(f"Unsafe repository path configured for {artifact.name}.") | |
| return artifact | |
| def _default_path(self, artifact: ModelArtifact) -> Path: | |
| return self.models_root.joinpath(*PurePosixPath(artifact.repository_path).parts) | |
| def _verify_existing(path: Path, artifact: ModelArtifact) -> Optional[Path]: | |
| if not path.exists(): | |
| return None | |
| if not path.is_file(): | |
| raise ModelArtifactFilesystemError( | |
| f"Expected {artifact.name} to be a regular file at {path}." | |
| ) | |
| actual = sha256_file(path) | |
| if actual != artifact.sha256: | |
| raise ModelArtifactChecksumError( | |
| f"Local {artifact.name} failed SHA-256 verification; expected {artifact.sha256}, received {actual}." | |
| ) | |
| return path | |
| def _download_function(self) -> DownloadFunction: | |
| if self.downloader is not None: | |
| return self.downloader | |
| try: | |
| from huggingface_hub import hf_hub_download | |
| except ImportError: | |
| raise ModelArtifactDownloadError( | |
| "huggingface_hub is required to retrieve missing private model artifacts." | |
| ) from None | |
| return hf_hub_download | |
| def resolve( | |
| self, | |
| key: str, | |
| *, | |
| local_path: Optional[Path] = None, | |
| download_if_missing: bool = True, | |
| ) -> Path: | |
| artifact = self._artifact(key) | |
| target = (local_path or self._default_path(artifact)).expanduser().resolve() | |
| verified = self._verify_existing(target, artifact) | |
| if verified is not None: | |
| return verified | |
| if not download_if_missing: | |
| raise ModelArtifactMissingError( | |
| f"Explicit checkpoint for {artifact.name} is missing at {target}; automatic download is disabled for configured overrides." | |
| ) | |
| token = self.environment.get("HF_TOKEN", "").strip() | |
| if not token: | |
| raise ModelArtifactAuthenticationError( | |
| f"{artifact.name} is missing and HF_TOKEN is required to access the private model repository." | |
| ) | |
| try: | |
| target.parent.mkdir(parents=True, exist_ok=True) | |
| except OSError as error: | |
| raise ModelArtifactFilesystemError( | |
| f"Could not create the local directory for {artifact.name}." | |
| ) from error | |
| lock_path = target.parent / f".{target.name}.lock" | |
| with _exclusive_file_lock(lock_path): | |
| verified = self._verify_existing(target, artifact) | |
| if verified is not None: | |
| return verified | |
| downloader = self._download_function() | |
| try: | |
| downloaded = Path( | |
| downloader( | |
| repo_id=self.repository_id, | |
| filename=artifact.repository_path, | |
| token=token, | |
| ) | |
| ).resolve() | |
| except Exception: | |
| raise ModelArtifactDownloadError( | |
| f"Could not retrieve {artifact.name} from the private model repository; verify HF_TOKEN access and network availability." | |
| ) from None | |
| if not downloaded.is_file(): | |
| raise ModelArtifactDownloadError( | |
| f"The private model repository did not return a file for {artifact.name}." | |
| ) | |
| actual = sha256_file(downloaded) | |
| if actual != artifact.sha256: | |
| raise ModelArtifactChecksumError( | |
| f"Downloaded {artifact.name} failed SHA-256 verification; expected {artifact.sha256}, received {actual}." | |
| ) | |
| temporary = target.parent / ( | |
| f".{target.name}.{os.getpid()}.{threading.get_ident()}.part" | |
| ) | |
| try: | |
| shutil.copyfile(downloaded, temporary) | |
| if sha256_file(temporary) != artifact.sha256: | |
| raise ModelArtifactChecksumError( | |
| f"Staged {artifact.name} failed SHA-256 verification." | |
| ) | |
| os.replace(temporary, target) | |
| except ModelArtifactError: | |
| raise | |
| except OSError as error: | |
| raise ModelArtifactFilesystemError( | |
| f"Could not publish the verified {artifact.name} to its local model path." | |
| ) from error | |
| finally: | |
| try: | |
| temporary.unlink(missing_ok=True) | |
| except OSError: | |
| pass | |
| return target | |
| _DEFAULT_RESOLVER = ModelArtifactResolver() | |
| def resolve_model_artifact( | |
| key: str, | |
| *, | |
| local_path: Optional[Path] = None, | |
| download_if_missing: bool = True, | |
| ) -> Path: | |
| """Resolve one of the pinned CodeBlueX artifacts.""" | |
| return _DEFAULT_RESOLVER.resolve( | |
| key, | |
| local_path=local_path, | |
| download_if_missing=download_if_missing, | |
| ) | |
| def resolve_configured_model_artifact( | |
| key: str, | |
| *, | |
| environment_key: str, | |
| local_path: Path, | |
| ) -> Path: | |
| """Honor an explicit checkpoint path without replacing a missing override.""" | |
| explicitly_configured = bool( | |
| _DEFAULT_RESOLVER.environment.get(environment_key, "").strip() | |
| ) | |
| return resolve_model_artifact( | |
| key, | |
| local_path=local_path, | |
| download_if_missing=not explicitly_configured, | |
| ) | |