geosave-engine 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.
- geosave_engine/__about__.py +2 -0
- geosave_engine/__init__.py +0 -0
- geosave_engine/cli/__init__.py +0 -0
- geosave_engine/cli/commands/__init__.py +1 -0
- geosave_engine/cli/commands/create.py +60 -0
- geosave_engine/cli/commands/infra.py +79 -0
- geosave_engine/cli/commands/upload.py +274 -0
- geosave_engine/cli/errors.py +12 -0
- geosave_engine/cli/main.py +11 -0
- geosave_engine/cli/workspace/__init__.py +12 -0
- geosave_engine/cli/workspace/artifact.py +186 -0
- geosave_engine/cli/workspace/model.py +125 -0
- geosave_engine/cli/workspace/scaffold.py +31 -0
- geosave_engine/cli/workspace/templates.py +37 -0
- geosave_engine/geodata/__init__.py +0 -0
- geosave_engine/geodata/datasets/__init__.py +22 -0
- geosave_engine/geodata/datasets/base_dataset.py +184 -0
- geosave_engine/geodata/datasets/coco_dataset.py +46 -0
- geosave_engine/geodata/datasets/geo_dataset.py +106 -0
- geosave_engine/geodata/datasets/geostack_dataset.py +127 -0
- geosave_engine/geodata/datasets/intersection_dataset.py +80 -0
- geosave_engine/geodata/datasets/non_geo_dataset.py +98 -0
- geosave_engine/geodata/datasets/samplers.py +34 -0
- geosave_engine/geodata/datasets/table_dataset.py +109 -0
- geosave_engine/geodata/datasets/yolo_dataset.py +108 -0
- geosave_engine/geodata/errors/__init__.py +5 -0
- geosave_engine/geodata/errors/errors.py +10 -0
- geosave_engine/geodata/features/__init__.py +45 -0
- geosave_engine/geodata/features/cloud_mask.py +116 -0
- geosave_engine/geodata/features/shadow_mask.py +45 -0
- geosave_engine/geodata/features/spectral_indices.py +350 -0
- geosave_engine/geodata/pipeline/__init__.py +21 -0
- geosave_engine/geodata/pipeline/anchor_sources.py +172 -0
- geosave_engine/geodata/pipeline/geo_pipeline.py +205 -0
- geosave_engine/geodata/sensors/__init__.py +19 -0
- geosave_engine/geodata/sensors/sensors.py +159 -0
- geosave_engine/geodata/sensors/sensors.yaml +95 -0
- geosave_engine/geodata/stac/__init__.py +7 -0
- geosave_engine/geodata/stac/client.py +138 -0
- geosave_engine/geodata/stac/query.py +133 -0
- geosave_engine/geodata/stac/source.py +488 -0
- geosave_engine/geodata/tile/__init__.py +26 -0
- geosave_engine/geodata/tile/geoanchor.py +452 -0
- geosave_engine/geodata/tile/geostack.py +237 -0
- geosave_engine/geodata/tile/geotile.py +449 -0
- geosave_engine/geodata/tile/ops.py +226 -0
- geosave_engine/geodata/utils/__init__.py +21 -0
- geosave_engine/geodata/utils/archives.py +41 -0
- geosave_engine/geodata/utils/crs.py +55 -0
- geosave_engine/geodata/utils/datetime.py +166 -0
- geosave_engine/geodata/utils/geodata.py +191 -0
- geosave_engine/geodata/utils/geolocator.py +121 -0
- geosave_engine/geodata/utils/geovis.py +574 -0
- geosave_engine/geodata/utils/io.py +164 -0
- geosave_engine/infra/.env.example +17 -0
- geosave_engine/infra/docker-compose.yml +191 -0
- geosave_engine/ml/__init__.py +0 -0
- geosave_engine/ml/callbacks/__init__.py +5 -0
- geosave_engine/ml/callbacks/prediction_logger.py +137 -0
- geosave_engine/ml/callbacks/prediction_writer.py +269 -0
- geosave_engine/ml/callbacks/threshold_calibrator.py +191 -0
- geosave_engine/ml/cli/__init__.py +3 -0
- geosave_engine/ml/cli/cli.py +143 -0
- geosave_engine/ml/inference/__init__.py +0 -0
- geosave_engine/ml/inference/protocol.py +45 -0
- geosave_engine/ml/inference/sliding_window.py +139 -0
- geosave_engine/ml/inference/thresholding.py +45 -0
- geosave_engine/ml/loss/__init__.py +3 -0
- geosave_engine/ml/loss/ohem.py +57 -0
- geosave_engine/ml/metrics/__init__.py +0 -0
- geosave_engine/ml/metrics/semantic_segmentation.py +157 -0
- geosave_engine/ml/models/__init__.py +33 -0
- geosave_engine/ml/models/contract/__init__.py +9 -0
- geosave_engine/ml/models/contract/chain.py +346 -0
- geosave_engine/ml/models/contract/context.py +269 -0
- geosave_engine/ml/models/contract/normalization.py +7 -0
- geosave_engine/ml/models/decoder/__init__.py +4 -0
- geosave_engine/ml/models/decoder/dpt.py +262 -0
- geosave_engine/ml/models/decoder/unet.py +224 -0
- geosave_engine/ml/models/encoder/__init__.py +5 -0
- geosave_engine/ml/models/encoder/clay.py +364 -0
- geosave_engine/ml/models/encoder/dinov3.py +173 -0
- geosave_engine/ml/models/encoder/prithvi.py +287 -0
- geosave_engine/ml/models/head/__init__.py +3 -0
- geosave_engine/ml/models/head/dense.py +91 -0
- geosave_engine/ml/models/monolith/__init__.py +3 -0
- geosave_engine/ml/models/monolith/ibm_granite_biomass.py +104 -0
- geosave_engine/ml/optimizer/__init__.py +0 -0
- geosave_engine/ml/optimizer/adagrad.py +27 -0
- geosave_engine/ml/optimizer/adam.py +41 -0
- geosave_engine/ml/optimizer/adamw.py +66 -0
- geosave_engine/ml/optimizer/rmsprop.py +44 -0
- geosave_engine/ml/optimizer/sgd.py +60 -0
- geosave_engine/ml/registry/__init__.py +15 -0
- geosave_engine/ml/registry/base.py +63 -0
- geosave_engine/ml/registry/loss.py +29 -0
- geosave_engine/ml/registry/model.py +174 -0
- geosave_engine/ml/registry/optimizer.py +38 -0
- geosave_engine/ml/registry/scheduler.py +32 -0
- geosave_engine/ml/tasks/__init__.py +3 -0
- geosave_engine/ml/tasks/semantic_segmentation.py +610 -0
- geosave_engine/ml/transforms/__init__.py +4 -0
- geosave_engine/ml/transforms/augmenter.py +83 -0
- geosave_engine/ml/transforms/processor.py +77 -0
- geosave_engine/ml/utils/__init__.py +11 -0
- geosave_engine/ml/utils/torch_params.py +81 -0
- geosave_engine/ml/utils/weights.py +31 -0
- geosave_engine/templates/common/.env +8 -0
- geosave_engine/templates/common/main.py +8 -0
- geosave_engine/templates/semantic_segmentation/supervised/configs/augmentation.yaml +7 -0
- geosave_engine/templates/semantic_segmentation/supervised/configs/metadata.yaml +9 -0
- geosave_engine/templates/semantic_segmentation/supervised/configs/model.yaml +30 -0
- geosave_engine/templates/semantic_segmentation/supervised/modules/data_pipeline.py +43 -0
- geosave_engine/utils/__init__.py +9 -0
- geosave_engine/utils/colorize.py +49 -0
- geosave_engine/utils/file_ops.py +47 -0
- geosave_engine/utils/fn.py +22 -0
- geosave_engine-0.1.0.dist-info/METADATA +159 -0
- geosave_engine-0.1.0.dist-info/RECORD +121 -0
- geosave_engine-0.1.0.dist-info/WHEEL +4 -0
- geosave_engine-0.1.0.dist-info/entry_points.txt +2 -0
|
File without changes
|
|
File without changes
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""GeoSave CLI command modules."""
|
|
@@ -0,0 +1,60 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from pathlib import Path
|
|
4
|
+
|
|
5
|
+
import questionary
|
|
6
|
+
import typer
|
|
7
|
+
|
|
8
|
+
from geosave_engine.cli.errors import AbortedByUserError
|
|
9
|
+
from geosave_engine.cli.workspace import Workspace, WorkspaceSpec
|
|
10
|
+
from geosave_engine.cli.workspace.templates import get_method_templates
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def create(
|
|
14
|
+
directory: Path = typer.Option(
|
|
15
|
+
Path("."),
|
|
16
|
+
"--dir",
|
|
17
|
+
"-d",
|
|
18
|
+
help="Directory to build the GeoSave workspace in.",
|
|
19
|
+
),
|
|
20
|
+
) -> None:
|
|
21
|
+
"""Create one GeoSave workspace."""
|
|
22
|
+
workspace = Workspace(directory, _prompt_workspace_spec())
|
|
23
|
+
workspace.setup_workspace()
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def _prompt_workspace_spec() -> WorkspaceSpec:
|
|
27
|
+
method_templates = get_method_templates()
|
|
28
|
+
project_name = _ask_required_text("Enter the project name:", "Project name")
|
|
29
|
+
project_task = _ask_required_choice(
|
|
30
|
+
"Select the main task for the project:",
|
|
31
|
+
list(method_templates),
|
|
32
|
+
"Project task",
|
|
33
|
+
)
|
|
34
|
+
project_method = _ask_required_choice(
|
|
35
|
+
"Select the method for the project:",
|
|
36
|
+
list(method_templates[project_task]),
|
|
37
|
+
"Project method",
|
|
38
|
+
)
|
|
39
|
+
description = questionary.text("Enter a description for the project (optional):").ask()
|
|
40
|
+
|
|
41
|
+
return WorkspaceSpec(
|
|
42
|
+
project_name=project_name,
|
|
43
|
+
project_task=project_task,
|
|
44
|
+
project_method=project_method,
|
|
45
|
+
description=description.strip() if description else None,
|
|
46
|
+
)
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _ask_required_text(question: str, field_name: str) -> str:
|
|
50
|
+
answer = questionary.text(question).ask()
|
|
51
|
+
if not answer or not answer.strip():
|
|
52
|
+
raise AbortedByUserError(f"{field_name} is required.")
|
|
53
|
+
return answer.strip()
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def _ask_required_choice(question: str, choices: list[str], field_name: str) -> str:
|
|
57
|
+
answer = questionary.select(question, choices=choices).ask()
|
|
58
|
+
if answer is None:
|
|
59
|
+
raise AbortedByUserError(f"{field_name} selection was aborted.")
|
|
60
|
+
return answer.strip()
|
|
@@ -0,0 +1,79 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import importlib.resources as pkg_resources
|
|
4
|
+
import shutil
|
|
5
|
+
import subprocess
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
|
|
8
|
+
import typer
|
|
9
|
+
|
|
10
|
+
infra_app = typer.Typer(help="Manage GeoSave dev infrastructure (Docker Compose).")
|
|
11
|
+
|
|
12
|
+
_COMPOSE_NAME = "docker-compose.yml"
|
|
13
|
+
_ENV_EXAMPLE = ".env.example"
|
|
14
|
+
_LOCAL_INFRA_DIR = "docker"
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def _bundled(filename: str) -> Path:
|
|
18
|
+
return Path(str(pkg_resources.files("geosave_engine") / "infra" / filename))
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def _local_compose() -> Path | None:
|
|
22
|
+
path = Path.cwd() / _LOCAL_INFRA_DIR / _COMPOSE_NAME
|
|
23
|
+
return path if path.exists() else None
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
@infra_app.command()
|
|
27
|
+
def init() -> None:
|
|
28
|
+
"""Copy docker-compose.yml and .env to ./docker/."""
|
|
29
|
+
destination = Path.cwd() / _LOCAL_INFRA_DIR
|
|
30
|
+
destination.mkdir(exist_ok=True, parents=True)
|
|
31
|
+
compose_destination = destination / _COMPOSE_NAME
|
|
32
|
+
env_destination = destination / ".env"
|
|
33
|
+
|
|
34
|
+
if compose_destination.exists():
|
|
35
|
+
typer.echo(f"{_COMPOSE_NAME} already exists — skipped.")
|
|
36
|
+
else:
|
|
37
|
+
shutil.copy2(_bundled(_COMPOSE_NAME), compose_destination)
|
|
38
|
+
typer.echo(f"Created docker/{_COMPOSE_NAME}")
|
|
39
|
+
|
|
40
|
+
if env_destination.exists():
|
|
41
|
+
typer.echo(".env already exists — skipped.")
|
|
42
|
+
else:
|
|
43
|
+
shutil.copy2(_bundled(_ENV_EXAMPLE), env_destination)
|
|
44
|
+
typer.echo("Created docker/.env from defaults")
|
|
45
|
+
|
|
46
|
+
typer.echo("\nNext: geosave infra up")
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
@infra_app.command()
|
|
50
|
+
def up(
|
|
51
|
+
profile: list[str] = typer.Option([], "--profile", "-p", help="Enable a service profile."),
|
|
52
|
+
detach: bool = typer.Option(
|
|
53
|
+
True,
|
|
54
|
+
"--detach/--no-detach",
|
|
55
|
+
"-d/-D",
|
|
56
|
+
help="Run in background.",
|
|
57
|
+
),
|
|
58
|
+
) -> None:
|
|
59
|
+
"""Start infrastructure containers."""
|
|
60
|
+
_run_compose(["up", *(["-d"] if detach else [])], profile)
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
@infra_app.command()
|
|
64
|
+
def down(profile: list[str] = typer.Option([], "--profile", "-p")) -> None:
|
|
65
|
+
"""Stop infrastructure containers."""
|
|
66
|
+
_run_compose(["down"], profile)
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
@infra_app.command()
|
|
70
|
+
def status() -> None:
|
|
71
|
+
"""Show running infrastructure containers."""
|
|
72
|
+
_run_compose(["ps"], [])
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def _run_compose(args: list[str], profiles: list[str]) -> None:
|
|
76
|
+
profile_flags = [flag for profile in profiles for flag in ("--profile", profile)]
|
|
77
|
+
compose_path = _local_compose() or _bundled(_COMPOSE_NAME)
|
|
78
|
+
command = ["docker", "compose", "-f", str(compose_path), *profile_flags, *args]
|
|
79
|
+
subprocess.run(command, check=True)
|
|
@@ -0,0 +1,274 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import importlib
|
|
4
|
+
import os
|
|
5
|
+
import pickle
|
|
6
|
+
import sys
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
from typing import TYPE_CHECKING
|
|
9
|
+
|
|
10
|
+
import questionary
|
|
11
|
+
import typer
|
|
12
|
+
import yaml
|
|
13
|
+
|
|
14
|
+
from geosave_engine.cli.errors import AbortedByUserError, WorkspaceError
|
|
15
|
+
|
|
16
|
+
if TYPE_CHECKING:
|
|
17
|
+
# Only ever used as type annotations (lazy strings, thanks to
|
|
18
|
+
# `from __future__ import annotations`) — never called by name at
|
|
19
|
+
# runtime, so keep them out of the module's real import graph. Both
|
|
20
|
+
# drag in ~10s of lightning.pytorch/mlflow import time, which main.py
|
|
21
|
+
# would otherwise pay on every `geosave` invocation (create, artifact,
|
|
22
|
+
# --help, ...), not just `upload`.
|
|
23
|
+
from lightning.pytorch import LightningModule
|
|
24
|
+
from mlflow.models.model import ModelInfo
|
|
25
|
+
from geosave_engine.cli.workspace import Workspace, load_run_artifact
|
|
26
|
+
from geosave_engine.cli.workspace.artifact import (
|
|
27
|
+
artifact_paths,
|
|
28
|
+
resolve_artifact_name,
|
|
29
|
+
select_checkpoint,
|
|
30
|
+
)
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def upload(
|
|
34
|
+
project_dir: Path | None = typer.Argument(None, help="Workspace directory."),
|
|
35
|
+
artifact_name: str | None = typer.Option(
|
|
36
|
+
None,
|
|
37
|
+
"--artifact",
|
|
38
|
+
"-a",
|
|
39
|
+
help="Artifact directory (for example, model_name/version_0).",
|
|
40
|
+
),
|
|
41
|
+
checkpoint: str | None = typer.Option(
|
|
42
|
+
None,
|
|
43
|
+
"--checkpoint",
|
|
44
|
+
"-c",
|
|
45
|
+
help="Checkpoint filename inside the run's checkpoints/ dir. Prompt if omitted and multiple exist.",
|
|
46
|
+
),
|
|
47
|
+
registered_model_name: str | None = typer.Option(
|
|
48
|
+
None,
|
|
49
|
+
"--name",
|
|
50
|
+
"-n",
|
|
51
|
+
help="MLflow registered model name. Defaults to the resolved run name.",
|
|
52
|
+
),
|
|
53
|
+
) -> None:
|
|
54
|
+
"""Upload one trained checkpoint to the MLflow model registry.
|
|
55
|
+
|
|
56
|
+
Resolve one run artifact (discover it under artifacts/, prompt if
|
|
57
|
+
ambiguous), rebuild its LightningModule from saved config + checkpoint
|
|
58
|
+
weights, then
|
|
59
|
+
log + register it to MLflow inside a fresh mlflow.start_run(), with
|
|
60
|
+
workspace modules/ bundled as code.
|
|
61
|
+
|
|
62
|
+
Args:
|
|
63
|
+
project_dir: Workspace directory. Defaults to cwd.
|
|
64
|
+
artifact_name: Artifact directory relative to artifacts/, for
|
|
65
|
+
example "DynamicWorld/version_9". Prompt if omitted.
|
|
66
|
+
checkpoint: Checkpoint filename. Prompt if omitted and multiple
|
|
67
|
+
checkpoints exist.
|
|
68
|
+
registered_model_name: MLflow registered model name. Defaults to
|
|
69
|
+
the resolved RunArtifact.model_name.
|
|
70
|
+
|
|
71
|
+
Raises:
|
|
72
|
+
WorkspaceError: If artifact/config is missing or invalid.
|
|
73
|
+
AbortedByUserError: If prompted for a missing MLFLOW_TRACKING_URI or
|
|
74
|
+
MLFLOW_EXPERIMENT_NAME and no answer is given.
|
|
75
|
+
"""
|
|
76
|
+
workspace = Workspace.load_workspace(project_dir or Path.cwd())
|
|
77
|
+
|
|
78
|
+
# modules/ has no __init__.py (namespace package) — only resolves once
|
|
79
|
+
# workspace.root is searchable. python main.py gets this for free from
|
|
80
|
+
# cwd; this installed console script doesn't, so make it explicit.
|
|
81
|
+
# No-op for premade (Path A) classes — those come from the installed
|
|
82
|
+
# geosave_engine package and resolve regardless.
|
|
83
|
+
sys.path.insert(0, str(workspace.root))
|
|
84
|
+
|
|
85
|
+
resolved_name = resolve_artifact_name(workspace, artifact_name)
|
|
86
|
+
run_dir = artifact_paths(workspace)[resolved_name]
|
|
87
|
+
artifact = load_run_artifact(run_dir)
|
|
88
|
+
|
|
89
|
+
checkpoint_path = select_checkpoint(artifact, checkpoint)
|
|
90
|
+
model = _load_model_from_checkpoint(artifact.config_path, checkpoint_path)
|
|
91
|
+
|
|
92
|
+
resolved_registered_name = registered_model_name or artifact.model_name
|
|
93
|
+
tracking_uri = _require_tracking_uri()
|
|
94
|
+
experiment_name = _require_experiment_name(resolved_registered_name)
|
|
95
|
+
|
|
96
|
+
model_info = log_model(
|
|
97
|
+
model=model,
|
|
98
|
+
name=resolved_registered_name,
|
|
99
|
+
checkpoint_path=checkpoint_path,
|
|
100
|
+
modules_dir=workspace.modules_dir,
|
|
101
|
+
tracking_uri=tracking_uri,
|
|
102
|
+
experiment_name=experiment_name,
|
|
103
|
+
)
|
|
104
|
+
# model_info.model_uri is mlflow 3.x's models:/m-<hash> model-id form —
|
|
105
|
+
# print the conventional name/version form instead, since that's what
|
|
106
|
+
# litserve configs and humans actually reference.
|
|
107
|
+
typer.echo(f"models:/{resolved_registered_name}/{model_info.registered_model_version}")
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
def _import_model_class(class_path: str) -> type[LightningModule]:
|
|
111
|
+
"""Dynamically import a LightningModule class from its dotted path.
|
|
112
|
+
|
|
113
|
+
Args:
|
|
114
|
+
class_path: Fully qualified class path (for example,
|
|
115
|
+
"geosave_engine.ml.tasks.SemanticSegmentationTask"), read from
|
|
116
|
+
the run's config.yaml "model.class_path".
|
|
117
|
+
|
|
118
|
+
Returns:
|
|
119
|
+
Imported class, ready for load_from_checkpoint.
|
|
120
|
+
|
|
121
|
+
Raises:
|
|
122
|
+
WorkspaceError: If the module or class can't be imported.
|
|
123
|
+
"""
|
|
124
|
+
module_name, _, class_name = class_path.rpartition(".")
|
|
125
|
+
try:
|
|
126
|
+
module = importlib.import_module(module_name)
|
|
127
|
+
return getattr(module, class_name)
|
|
128
|
+
except (ImportError, AttributeError) as error:
|
|
129
|
+
raise WorkspaceError(f"Could not import model class '{class_path}': {error}") from error
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
def _load_model_from_checkpoint(config_path: Path, checkpoint_path: Path) -> LightningModule:
|
|
133
|
+
"""Rebuild one trained LightningModule from its run artifacts.
|
|
134
|
+
|
|
135
|
+
Uses the checkpoint's saved hyperparameters (save_hyperparameters() at
|
|
136
|
+
training time) — no need to re-pass init_args from config.yaml.
|
|
137
|
+
|
|
138
|
+
Args:
|
|
139
|
+
config_path: Run's config.yaml, used only to resolve the model
|
|
140
|
+
class_path.
|
|
141
|
+
checkpoint_path: Checkpoint file with saved weights + hparams.
|
|
142
|
+
|
|
143
|
+
Returns:
|
|
144
|
+
Instantiated model with weights loaded, ready to log.
|
|
145
|
+
|
|
146
|
+
Raises:
|
|
147
|
+
WorkspaceError: If config.yaml lacks a "model.class_path" entry, or
|
|
148
|
+
the checkpoint file is unreadable or corrupted.
|
|
149
|
+
"""
|
|
150
|
+
from geosave_engine.ml.inference.protocol import Predictable
|
|
151
|
+
|
|
152
|
+
config = yaml.safe_load(config_path.read_text())
|
|
153
|
+
class_path = config.get("model", {}).get("class_path")
|
|
154
|
+
if not class_path:
|
|
155
|
+
raise WorkspaceError(f"config.yaml missing model.class_path: {config_path}")
|
|
156
|
+
|
|
157
|
+
model_cls = _import_model_class(class_path)
|
|
158
|
+
try:
|
|
159
|
+
model = model_cls.load_from_checkpoint(str(checkpoint_path), map_location="cpu")
|
|
160
|
+
except (RuntimeError, OSError, EOFError, pickle.UnpicklingError) as error:
|
|
161
|
+
raise WorkspaceError(f"Could not load checkpoint '{checkpoint_path.name}': {error}") from error
|
|
162
|
+
|
|
163
|
+
if not isinstance(model, Predictable):
|
|
164
|
+
raise WorkspaceError(
|
|
165
|
+
f"{class_path} doesn't implement Predictable (no usable predict() method) — "
|
|
166
|
+
"can't register a model that can't be served."
|
|
167
|
+
)
|
|
168
|
+
return model
|
|
169
|
+
|
|
170
|
+
|
|
171
|
+
def _require_tracking_uri() -> str:
|
|
172
|
+
"""Read MLFLOW_TRACKING_URI from the environment, prompting when unset.
|
|
173
|
+
|
|
174
|
+
No local-mlruns fallback — upload always targets a real registry, so an
|
|
175
|
+
empty answer aborts instead of silently defaulting.
|
|
176
|
+
|
|
177
|
+
Returns:
|
|
178
|
+
Tracking URI value.
|
|
179
|
+
|
|
180
|
+
Raises:
|
|
181
|
+
AbortedByUserError: If prompted and the user gives no answer.
|
|
182
|
+
"""
|
|
183
|
+
tracking_uri = os.getenv("MLFLOW_TRACKING_URI")
|
|
184
|
+
if tracking_uri:
|
|
185
|
+
return tracking_uri
|
|
186
|
+
|
|
187
|
+
answer = questionary.text(
|
|
188
|
+
"MLFLOW_TRACKING_URI is not set. Enter your MLflow tracking URI:"
|
|
189
|
+
).ask()
|
|
190
|
+
if not answer or not answer.strip():
|
|
191
|
+
raise AbortedByUserError("MLflow tracking URI is required.")
|
|
192
|
+
return answer.strip()
|
|
193
|
+
|
|
194
|
+
|
|
195
|
+
def _require_experiment_name(default: str) -> str:
|
|
196
|
+
"""Read MLFLOW_EXPERIMENT_NAME from the environment, prompting when unset.
|
|
197
|
+
|
|
198
|
+
Mirrors GeosaveCLI's own training-time default (falls back to
|
|
199
|
+
model_name when MLFLOW_EXPERIMENT_NAME isn't set — see ml/cli/cli.py).
|
|
200
|
+
|
|
201
|
+
Args:
|
|
202
|
+
default: Prefilled answer, typically the resolved model name.
|
|
203
|
+
|
|
204
|
+
Returns:
|
|
205
|
+
Experiment name value.
|
|
206
|
+
|
|
207
|
+
Raises:
|
|
208
|
+
AbortedByUserError: If prompted and the user gives no answer.
|
|
209
|
+
"""
|
|
210
|
+
experiment_name = os.getenv("MLFLOW_EXPERIMENT_NAME")
|
|
211
|
+
if experiment_name:
|
|
212
|
+
return experiment_name
|
|
213
|
+
|
|
214
|
+
answer = questionary.text(
|
|
215
|
+
"MLFLOW_EXPERIMENT_NAME is not set. Enter an experiment name:",
|
|
216
|
+
default=default,
|
|
217
|
+
).ask()
|
|
218
|
+
if not answer or not answer.strip():
|
|
219
|
+
raise AbortedByUserError("MLflow experiment name is required.")
|
|
220
|
+
return answer.strip()
|
|
221
|
+
|
|
222
|
+
|
|
223
|
+
def log_model(
|
|
224
|
+
model: LightningModule,
|
|
225
|
+
name: str,
|
|
226
|
+
checkpoint_path: Path,
|
|
227
|
+
modules_dir: Path,
|
|
228
|
+
tracking_uri: str,
|
|
229
|
+
experiment_name: str,
|
|
230
|
+
) -> ModelInfo:
|
|
231
|
+
"""Log one model run to MLflow, resolving tracking uri/experiment/run itself.
|
|
232
|
+
|
|
233
|
+
MLflow acts as registry using pytorch log_model. modules_dir is
|
|
234
|
+
bundled as code_paths so registry-side consumers (litserve) can import
|
|
235
|
+
workspace-local preprocessing (for example, modules.data_pipeline) after
|
|
236
|
+
pulling the model; litserve owns preprocessing + inference wiring, not
|
|
237
|
+
the logged artifact.
|
|
238
|
+
|
|
239
|
+
Args:
|
|
240
|
+
model: Instantiated model to upload.
|
|
241
|
+
name: Resolved name (for example, "DynamicWorld") — used as the
|
|
242
|
+
MLflow run name, the artifact's name inside it, and the
|
|
243
|
+
registered model name. upload only exposes one --name flag, so
|
|
244
|
+
these three never differ in practice.
|
|
245
|
+
checkpoint_path: Local checkpoint this model was rebuilt from,
|
|
246
|
+
stored as metadata for traceability.
|
|
247
|
+
modules_dir: Workspace modules/ directory bundled as code_paths —
|
|
248
|
+
code MLflow copies into the model's code/ dir and prepends to
|
|
249
|
+
sys.path on load.
|
|
250
|
+
tracking_uri: MLflow tracking URI to log against.
|
|
251
|
+
experiment_name: MLflow experiment to log under.
|
|
252
|
+
|
|
253
|
+
Returns:
|
|
254
|
+
Logged MLflow model information.
|
|
255
|
+
"""
|
|
256
|
+
import mlflow
|
|
257
|
+
|
|
258
|
+
mlflow.set_tracking_uri(tracking_uri)
|
|
259
|
+
mlflow.set_experiment(experiment_name)
|
|
260
|
+
|
|
261
|
+
typer.echo(f"Uploading '{name}' ({checkpoint_path.name}) to {tracking_uri} ...", err=True)
|
|
262
|
+
with mlflow.start_run(run_name=name):
|
|
263
|
+
# MLFLOW_ENABLE_ARTIFACTS_PROGRESS_BAR defaults on — mlflow shows
|
|
264
|
+
# its own tqdm byte-level progress for the artifact upload itself.
|
|
265
|
+
model_info = mlflow.pytorch.log_model(
|
|
266
|
+
pytorch_model=model,
|
|
267
|
+
name=name,
|
|
268
|
+
code_paths=[str(modules_dir)],
|
|
269
|
+
registered_model_name=name,
|
|
270
|
+
metadata={"checkpoint": checkpoint_path.name},
|
|
271
|
+
)
|
|
272
|
+
typer.echo(f"Registered '{name}' version {model_info.registered_model_version}", err=True)
|
|
273
|
+
|
|
274
|
+
return model_info
|
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
class GeosaveCliError(Exception):
|
|
2
|
+
"""Base error raised by GeoSave commands."""
|
|
3
|
+
|
|
4
|
+
exit_code: int = 1
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class AbortedByUserError(GeosaveCliError):
|
|
8
|
+
"""The user cancelled a prompt (Ctrl-C, empty answer)."""
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class WorkspaceError(GeosaveCliError):
|
|
12
|
+
"""A runtime command could not locate or load a workspace."""
|
|
@@ -0,0 +1,11 @@
|
|
|
1
|
+
import typer
|
|
2
|
+
|
|
3
|
+
from geosave_engine.cli.commands.create import create
|
|
4
|
+
from geosave_engine.cli.commands.infra import infra_app
|
|
5
|
+
from geosave_engine.cli.commands.upload import upload
|
|
6
|
+
|
|
7
|
+
app = typer.Typer(help="GeoSave Engine CLI", no_args_is_help=True)
|
|
8
|
+
app.command()(create)
|
|
9
|
+
app.command()(upload)
|
|
10
|
+
|
|
11
|
+
app.add_typer(infra_app, name="infra")
|
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
"""Workspace loading, discovery, and scaffolding."""
|
|
2
|
+
|
|
3
|
+
from .artifact import RunArtifact, discover_artifacts, load_run_artifact
|
|
4
|
+
from .model import Workspace, WorkspaceSpec
|
|
5
|
+
|
|
6
|
+
__all__ = [
|
|
7
|
+
"RunArtifact",
|
|
8
|
+
"Workspace",
|
|
9
|
+
"WorkspaceSpec",
|
|
10
|
+
"discover_artifacts",
|
|
11
|
+
"load_run_artifact",
|
|
12
|
+
]
|
|
@@ -0,0 +1,186 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
from pathlib import Path
|
|
5
|
+
from typing import TYPE_CHECKING
|
|
6
|
+
|
|
7
|
+
import questionary
|
|
8
|
+
|
|
9
|
+
from geosave_engine.cli.errors import AbortedByUserError, WorkspaceError
|
|
10
|
+
|
|
11
|
+
if TYPE_CHECKING:
|
|
12
|
+
from .model import Workspace
|
|
13
|
+
|
|
14
|
+
_CONFIG_FILE = "config.yaml"
|
|
15
|
+
_CHECKPOINTS_DIR = "checkpoints"
|
|
16
|
+
_CHECKPOINT_GLOB = "*.ckpt"
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
@dataclass(frozen=True)
|
|
20
|
+
class RunArtifact:
|
|
21
|
+
"""Store paths and identity for one model run.
|
|
22
|
+
|
|
23
|
+
Args:
|
|
24
|
+
run_dir: Version directory containing artifacts for one run
|
|
25
|
+
(for example, artifacts/DynamicWorld/version_9).
|
|
26
|
+
config_path: Lightning config saved for the run.
|
|
27
|
+
checkpoint_paths: All discovered checkpoints for the run, sorted.
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
run_dir: Path
|
|
31
|
+
config_path: Path
|
|
32
|
+
checkpoint_paths: list[Path]
|
|
33
|
+
|
|
34
|
+
@property
|
|
35
|
+
def model_name(self) -> str:
|
|
36
|
+
"""Return the run's parent directory name (for example, "DynamicWorld")."""
|
|
37
|
+
return self.run_dir.parent.name
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def discover_artifacts(artifacts_dir: Path) -> list[Path]:
|
|
41
|
+
"""Find run version directories without reading their contents.
|
|
42
|
+
|
|
43
|
+
Args:
|
|
44
|
+
artifacts_dir: Workspace artifacts directory.
|
|
45
|
+
|
|
46
|
+
Returns:
|
|
47
|
+
Sorted version directories (artifacts/<model_name>/version_N) that
|
|
48
|
+
hold a config.yaml.
|
|
49
|
+
"""
|
|
50
|
+
if not artifacts_dir.is_dir():
|
|
51
|
+
return []
|
|
52
|
+
|
|
53
|
+
return sorted(
|
|
54
|
+
path.resolve()
|
|
55
|
+
for path in artifacts_dir.glob("*/version_*")
|
|
56
|
+
if path.is_dir() and (path / _CONFIG_FILE).is_file()
|
|
57
|
+
)
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def load_run_artifact(run_dir: Path) -> RunArtifact:
|
|
61
|
+
"""Load one run directory using the canonical artifact layout.
|
|
62
|
+
|
|
63
|
+
Args:
|
|
64
|
+
run_dir: Version directory containing one model run.
|
|
65
|
+
|
|
66
|
+
Returns:
|
|
67
|
+
Validated paths for the run.
|
|
68
|
+
|
|
69
|
+
Raises:
|
|
70
|
+
WorkspaceError: If required files are missing.
|
|
71
|
+
"""
|
|
72
|
+
resolved_run_dir = run_dir.expanduser().resolve()
|
|
73
|
+
if not resolved_run_dir.is_dir():
|
|
74
|
+
raise WorkspaceError(f"Artifact run directory not found: {resolved_run_dir}")
|
|
75
|
+
|
|
76
|
+
config_path = _require_file(resolved_run_dir / _CONFIG_FILE)
|
|
77
|
+
checkpoint_paths = _discover_checkpoints(resolved_run_dir / _CHECKPOINTS_DIR)
|
|
78
|
+
|
|
79
|
+
return RunArtifact(
|
|
80
|
+
run_dir=resolved_run_dir,
|
|
81
|
+
config_path=config_path,
|
|
82
|
+
checkpoint_paths=checkpoint_paths,
|
|
83
|
+
)
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def _require_file(path: Path) -> Path:
|
|
87
|
+
if not path.is_file():
|
|
88
|
+
raise WorkspaceError(f"Required artifact file not found: {path}")
|
|
89
|
+
return path.resolve()
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def _discover_checkpoints(checkpoints_dir: Path) -> list[Path]:
|
|
93
|
+
if not checkpoints_dir.is_dir():
|
|
94
|
+
raise WorkspaceError(f"Checkpoints directory not found: {checkpoints_dir}")
|
|
95
|
+
|
|
96
|
+
checkpoint_paths = sorted(path.resolve() for path in checkpoints_dir.glob(_CHECKPOINT_GLOB))
|
|
97
|
+
if not checkpoint_paths:
|
|
98
|
+
raise WorkspaceError(f"No checkpoint files found in: {checkpoints_dir}")
|
|
99
|
+
return checkpoint_paths
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
def artifact_paths(workspace: Workspace) -> dict[str, Path]:
|
|
103
|
+
"""Map artifact keys to their run directories.
|
|
104
|
+
|
|
105
|
+
Args:
|
|
106
|
+
workspace: Loaded workspace to scan.
|
|
107
|
+
|
|
108
|
+
Returns:
|
|
109
|
+
Keys like "model_name/version_0" mapped to their run directory.
|
|
110
|
+
"""
|
|
111
|
+
return {
|
|
112
|
+
str(path.relative_to(workspace.artifacts_dir)): path
|
|
113
|
+
for path in workspace.artifacts
|
|
114
|
+
}
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
def resolve_artifact_name(workspace: Workspace, artifact_name: str | None) -> str:
|
|
118
|
+
"""Resolve one artifact key, prompting when omitted.
|
|
119
|
+
|
|
120
|
+
Args:
|
|
121
|
+
workspace: Loaded workspace to scan.
|
|
122
|
+
artifact_name: Explicit artifact key, skipping the prompt.
|
|
123
|
+
|
|
124
|
+
Returns:
|
|
125
|
+
Validated artifact key, usable with artifact_paths(workspace).
|
|
126
|
+
|
|
127
|
+
Raises:
|
|
128
|
+
WorkspaceError: If no artifacts exist, or artifact_name doesn't
|
|
129
|
+
match any.
|
|
130
|
+
AbortedByUserError: If the prompt is cancelled.
|
|
131
|
+
"""
|
|
132
|
+
paths = artifact_paths(workspace)
|
|
133
|
+
if artifact_name is None:
|
|
134
|
+
if not paths:
|
|
135
|
+
raise WorkspaceError(f"No artifacts found in: {workspace.artifacts_dir}")
|
|
136
|
+
artifact_name = _prompt_for_artifact(list(paths))
|
|
137
|
+
|
|
138
|
+
if artifact_name not in paths:
|
|
139
|
+
available = ", ".join(sorted(paths))
|
|
140
|
+
raise WorkspaceError(f"Artifact not found: {artifact_name}. Available: {available}")
|
|
141
|
+
|
|
142
|
+
return artifact_name
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
def select_checkpoint(artifact: RunArtifact, checkpoint: str | None) -> Path:
|
|
146
|
+
"""Resolve one checkpoint from a loaded RunArtifact, prompting if ambiguous.
|
|
147
|
+
|
|
148
|
+
Args:
|
|
149
|
+
artifact: Loaded run artifact with all discovered checkpoint_paths.
|
|
150
|
+
checkpoint: Checkpoint filename to use directly, skipping the
|
|
151
|
+
prompt. Prompt only when multiple checkpoints exist.
|
|
152
|
+
|
|
153
|
+
Returns:
|
|
154
|
+
Resolved checkpoint path.
|
|
155
|
+
|
|
156
|
+
Raises:
|
|
157
|
+
WorkspaceError: If checkpoint doesn't match any discovered file.
|
|
158
|
+
AbortedByUserError: If the prompt is cancelled.
|
|
159
|
+
"""
|
|
160
|
+
checkpoint_paths = artifact.checkpoint_paths
|
|
161
|
+
|
|
162
|
+
if checkpoint is not None:
|
|
163
|
+
matches = [path for path in checkpoint_paths if path.name == checkpoint]
|
|
164
|
+
if not matches:
|
|
165
|
+
available = ", ".join(path.name for path in checkpoint_paths)
|
|
166
|
+
raise WorkspaceError(f"Checkpoint not found: {checkpoint}. Available: {available}")
|
|
167
|
+
return matches[0]
|
|
168
|
+
|
|
169
|
+
if len(checkpoint_paths) == 1:
|
|
170
|
+
return checkpoint_paths[0]
|
|
171
|
+
|
|
172
|
+
answer = questionary.select(
|
|
173
|
+
"Select a checkpoint:",
|
|
174
|
+
choices=[path.name for path in checkpoint_paths],
|
|
175
|
+
).ask()
|
|
176
|
+
if answer is None:
|
|
177
|
+
raise AbortedByUserError("Checkpoint selection was aborted by the user.")
|
|
178
|
+
|
|
179
|
+
return next(path for path in checkpoint_paths if path.name == answer.strip())
|
|
180
|
+
|
|
181
|
+
|
|
182
|
+
def _prompt_for_artifact(artifact_keys: list[str]) -> str:
|
|
183
|
+
answer = questionary.select("Select an artifact:", choices=artifact_keys).ask()
|
|
184
|
+
if answer is None:
|
|
185
|
+
raise AbortedByUserError("Artifact selection was aborted by the user.")
|
|
186
|
+
return answer.strip()
|