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.
Files changed (121) hide show
  1. geosave_engine/__about__.py +2 -0
  2. geosave_engine/__init__.py +0 -0
  3. geosave_engine/cli/__init__.py +0 -0
  4. geosave_engine/cli/commands/__init__.py +1 -0
  5. geosave_engine/cli/commands/create.py +60 -0
  6. geosave_engine/cli/commands/infra.py +79 -0
  7. geosave_engine/cli/commands/upload.py +274 -0
  8. geosave_engine/cli/errors.py +12 -0
  9. geosave_engine/cli/main.py +11 -0
  10. geosave_engine/cli/workspace/__init__.py +12 -0
  11. geosave_engine/cli/workspace/artifact.py +186 -0
  12. geosave_engine/cli/workspace/model.py +125 -0
  13. geosave_engine/cli/workspace/scaffold.py +31 -0
  14. geosave_engine/cli/workspace/templates.py +37 -0
  15. geosave_engine/geodata/__init__.py +0 -0
  16. geosave_engine/geodata/datasets/__init__.py +22 -0
  17. geosave_engine/geodata/datasets/base_dataset.py +184 -0
  18. geosave_engine/geodata/datasets/coco_dataset.py +46 -0
  19. geosave_engine/geodata/datasets/geo_dataset.py +106 -0
  20. geosave_engine/geodata/datasets/geostack_dataset.py +127 -0
  21. geosave_engine/geodata/datasets/intersection_dataset.py +80 -0
  22. geosave_engine/geodata/datasets/non_geo_dataset.py +98 -0
  23. geosave_engine/geodata/datasets/samplers.py +34 -0
  24. geosave_engine/geodata/datasets/table_dataset.py +109 -0
  25. geosave_engine/geodata/datasets/yolo_dataset.py +108 -0
  26. geosave_engine/geodata/errors/__init__.py +5 -0
  27. geosave_engine/geodata/errors/errors.py +10 -0
  28. geosave_engine/geodata/features/__init__.py +45 -0
  29. geosave_engine/geodata/features/cloud_mask.py +116 -0
  30. geosave_engine/geodata/features/shadow_mask.py +45 -0
  31. geosave_engine/geodata/features/spectral_indices.py +350 -0
  32. geosave_engine/geodata/pipeline/__init__.py +21 -0
  33. geosave_engine/geodata/pipeline/anchor_sources.py +172 -0
  34. geosave_engine/geodata/pipeline/geo_pipeline.py +205 -0
  35. geosave_engine/geodata/sensors/__init__.py +19 -0
  36. geosave_engine/geodata/sensors/sensors.py +159 -0
  37. geosave_engine/geodata/sensors/sensors.yaml +95 -0
  38. geosave_engine/geodata/stac/__init__.py +7 -0
  39. geosave_engine/geodata/stac/client.py +138 -0
  40. geosave_engine/geodata/stac/query.py +133 -0
  41. geosave_engine/geodata/stac/source.py +488 -0
  42. geosave_engine/geodata/tile/__init__.py +26 -0
  43. geosave_engine/geodata/tile/geoanchor.py +452 -0
  44. geosave_engine/geodata/tile/geostack.py +237 -0
  45. geosave_engine/geodata/tile/geotile.py +449 -0
  46. geosave_engine/geodata/tile/ops.py +226 -0
  47. geosave_engine/geodata/utils/__init__.py +21 -0
  48. geosave_engine/geodata/utils/archives.py +41 -0
  49. geosave_engine/geodata/utils/crs.py +55 -0
  50. geosave_engine/geodata/utils/datetime.py +166 -0
  51. geosave_engine/geodata/utils/geodata.py +191 -0
  52. geosave_engine/geodata/utils/geolocator.py +121 -0
  53. geosave_engine/geodata/utils/geovis.py +574 -0
  54. geosave_engine/geodata/utils/io.py +164 -0
  55. geosave_engine/infra/.env.example +17 -0
  56. geosave_engine/infra/docker-compose.yml +191 -0
  57. geosave_engine/ml/__init__.py +0 -0
  58. geosave_engine/ml/callbacks/__init__.py +5 -0
  59. geosave_engine/ml/callbacks/prediction_logger.py +137 -0
  60. geosave_engine/ml/callbacks/prediction_writer.py +269 -0
  61. geosave_engine/ml/callbacks/threshold_calibrator.py +191 -0
  62. geosave_engine/ml/cli/__init__.py +3 -0
  63. geosave_engine/ml/cli/cli.py +143 -0
  64. geosave_engine/ml/inference/__init__.py +0 -0
  65. geosave_engine/ml/inference/protocol.py +45 -0
  66. geosave_engine/ml/inference/sliding_window.py +139 -0
  67. geosave_engine/ml/inference/thresholding.py +45 -0
  68. geosave_engine/ml/loss/__init__.py +3 -0
  69. geosave_engine/ml/loss/ohem.py +57 -0
  70. geosave_engine/ml/metrics/__init__.py +0 -0
  71. geosave_engine/ml/metrics/semantic_segmentation.py +157 -0
  72. geosave_engine/ml/models/__init__.py +33 -0
  73. geosave_engine/ml/models/contract/__init__.py +9 -0
  74. geosave_engine/ml/models/contract/chain.py +346 -0
  75. geosave_engine/ml/models/contract/context.py +269 -0
  76. geosave_engine/ml/models/contract/normalization.py +7 -0
  77. geosave_engine/ml/models/decoder/__init__.py +4 -0
  78. geosave_engine/ml/models/decoder/dpt.py +262 -0
  79. geosave_engine/ml/models/decoder/unet.py +224 -0
  80. geosave_engine/ml/models/encoder/__init__.py +5 -0
  81. geosave_engine/ml/models/encoder/clay.py +364 -0
  82. geosave_engine/ml/models/encoder/dinov3.py +173 -0
  83. geosave_engine/ml/models/encoder/prithvi.py +287 -0
  84. geosave_engine/ml/models/head/__init__.py +3 -0
  85. geosave_engine/ml/models/head/dense.py +91 -0
  86. geosave_engine/ml/models/monolith/__init__.py +3 -0
  87. geosave_engine/ml/models/monolith/ibm_granite_biomass.py +104 -0
  88. geosave_engine/ml/optimizer/__init__.py +0 -0
  89. geosave_engine/ml/optimizer/adagrad.py +27 -0
  90. geosave_engine/ml/optimizer/adam.py +41 -0
  91. geosave_engine/ml/optimizer/adamw.py +66 -0
  92. geosave_engine/ml/optimizer/rmsprop.py +44 -0
  93. geosave_engine/ml/optimizer/sgd.py +60 -0
  94. geosave_engine/ml/registry/__init__.py +15 -0
  95. geosave_engine/ml/registry/base.py +63 -0
  96. geosave_engine/ml/registry/loss.py +29 -0
  97. geosave_engine/ml/registry/model.py +174 -0
  98. geosave_engine/ml/registry/optimizer.py +38 -0
  99. geosave_engine/ml/registry/scheduler.py +32 -0
  100. geosave_engine/ml/tasks/__init__.py +3 -0
  101. geosave_engine/ml/tasks/semantic_segmentation.py +610 -0
  102. geosave_engine/ml/transforms/__init__.py +4 -0
  103. geosave_engine/ml/transforms/augmenter.py +83 -0
  104. geosave_engine/ml/transforms/processor.py +77 -0
  105. geosave_engine/ml/utils/__init__.py +11 -0
  106. geosave_engine/ml/utils/torch_params.py +81 -0
  107. geosave_engine/ml/utils/weights.py +31 -0
  108. geosave_engine/templates/common/.env +8 -0
  109. geosave_engine/templates/common/main.py +8 -0
  110. geosave_engine/templates/semantic_segmentation/supervised/configs/augmentation.yaml +7 -0
  111. geosave_engine/templates/semantic_segmentation/supervised/configs/metadata.yaml +9 -0
  112. geosave_engine/templates/semantic_segmentation/supervised/configs/model.yaml +30 -0
  113. geosave_engine/templates/semantic_segmentation/supervised/modules/data_pipeline.py +43 -0
  114. geosave_engine/utils/__init__.py +9 -0
  115. geosave_engine/utils/colorize.py +49 -0
  116. geosave_engine/utils/file_ops.py +47 -0
  117. geosave_engine/utils/fn.py +22 -0
  118. geosave_engine-0.1.0.dist-info/METADATA +159 -0
  119. geosave_engine-0.1.0.dist-info/RECORD +121 -0
  120. geosave_engine-0.1.0.dist-info/WHEEL +4 -0
  121. geosave_engine-0.1.0.dist-info/entry_points.txt +2 -0
@@ -0,0 +1,2 @@
1
+ __version__ = "0.1.0"
2
+ __author__ = "Fatah Muria"
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()