Skip to content

Commit f498682

Browse files
Clay inference model
1 parent 049dee3 commit f498682

6 files changed

Lines changed: 589 additions & 89 deletions

File tree

‎s2-clay/clay_inference.py‎

Lines changed: 289 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,289 @@
1+
import math
2+
import os
3+
from datetime import datetime
4+
from functools import lru_cache
5+
from pathlib import Path
6+
7+
import boto3
8+
import numpy as np
9+
import torch
10+
import xarray as xr
11+
import yaml
12+
import zarr
13+
from botocore.config import Config
14+
from box import Box
15+
from claymodel.module import ClayMAEModule
16+
from cyclopts import App
17+
from dotenv import load_dotenv
18+
from odc.geo.geobox import GeoBox
19+
from tilebox.workflows import Client as WorkflowsClient
20+
from tilebox.workflows import ExecutionContext, Task
21+
from tilebox.workflows.observability.logging import configure_console_logging, configure_otel_logging_axiom, get_logger
22+
from tilebox.workflows.observability.tracing import configure_otel_tracing_axiom
23+
from torchvision.transforms.v2 import Normalize, Transform
24+
25+
from sentinel2zarr import (
26+
COMPRESSOR,
27+
OUTPUT_BUCKET,
28+
Chunk2D,
29+
OTCBucketCache,
30+
RegionOfInterest,
31+
open_zarr_store,
32+
)
33+
34+
logger = get_logger()
35+
36+
CLAY_INFERENCE_TILE_SIZE = 256 # input tile size for the model is 256x256 pixels
37+
CLAY_PATCH_SIZE = 8 # the model computes embeddings for 8x8 patches within each tile
38+
CLAY_EMBEDDING_DIM = 1024 # embedding dimensionality of the model
39+
40+
# wget -q https://huggingface.co/made-with-clay/Clay/resolve/main/v1.5/clay-v1.5.ckpt
41+
_CLAY_CHECKPOINT = Path(__file__).parent / "clay-v1.5.ckpt"
42+
_CLAY_METADATA = Path(__file__).parent / "configs/metadata.yaml"
43+
_CLAY_PLATFORM = "sentinel-2-l2a"
44+
45+
46+
@lru_cache
47+
def device() -> torch.device:
48+
if torch.cuda.is_available():
49+
logger.info("CUDA is available, using GPU")
50+
return torch.device("cuda:0") # use the GPU if available
51+
if torch.backends.mps.is_available():
52+
logger.info("MPS is available, using Mac GPU")
53+
return torch.device("mps:0") # use the GPU if available
54+
logger.info("CUDA is not available, falling back to CPU")
55+
return torch.device("cpu") # otherwise fall back to CPU
56+
57+
58+
@lru_cache
59+
def clay_model() -> ClayMAEModule:
60+
"""Load the Clay model weights into memory"""
61+
logger.info("Loading Clay model weights into memory")
62+
model = ClayMAEModule.load_from_checkpoint(
63+
_CLAY_CHECKPOINT,
64+
model_size="large",
65+
metadata_path=_CLAY_METADATA.as_posix(),
66+
dolls=[16, 32, 64, 128, 256, 768, 1024],
67+
doll_weights=[1, 1, 1, 1, 1, 1, 1],
68+
mask_ratio=0.0,
69+
shuffle=False,
70+
)
71+
return model.to(device()).eval()
72+
73+
74+
class ClayInferenceOnMosaic(Task):
75+
mosaic_zarr_group: str
76+
"""Path to the zarr group containing the mosaic to run inference on. The group is expected to have a "mosaic" array
77+
with the shape (band, y, x) and a "band" array with the shape (band,) containing the band names as strings.
78+
"""
79+
80+
roi: RegionOfInterest
81+
"""The region of interest that the mosaic was computed for"""
82+
83+
crs: str
84+
"""The CRS of the mosaic"""
85+
86+
resolution: float
87+
"""The resolution of the mosaic in units of the CRS"""
88+
89+
output_zarr: tuple[str, str]
90+
"""The path to the output zarr group and the name of the output array"""
91+
92+
def execute(self, context: ExecutionContext) -> None:
93+
geobox = self.roi.area.as_geobox(self.crs, self.resolution)
94+
95+
output_group, output_array = self.output_zarr
96+
store = open_zarr_store(output_group)
97+
zarr.create_array(
98+
store=store,
99+
name=output_array,
100+
shape=(geobox.shape.y // CLAY_PATCH_SIZE, geobox.shape.x // CLAY_PATCH_SIZE, CLAY_EMBEDDING_DIM),
101+
chunks=(
102+
CLAY_INFERENCE_TILE_SIZE // CLAY_PATCH_SIZE, # 32
103+
CLAY_INFERENCE_TILE_SIZE // CLAY_PATCH_SIZE, # 32
104+
CLAY_EMBEDDING_DIM, # 1024
105+
),
106+
dimension_names=("y", "x", "embedding"),
107+
compressors=COMPRESSOR,
108+
dtype=np.float32,
109+
overwrite=True,
110+
)
111+
112+
chunks = self.roi.area.chunks(self.crs, self.resolution, (CLAY_INFERENCE_TILE_SIZE, CLAY_INFERENCE_TILE_SIZE))
113+
for chunk in chunks:
114+
context.submit_subtask(
115+
ClayInferenceTile(
116+
chunk,
117+
self.mosaic_zarr_group,
118+
self.roi,
119+
self.crs,
120+
self.resolution,
121+
self.output_zarr,
122+
),
123+
)
124+
context.progress("inference").add(len(chunks))
125+
126+
127+
@lru_cache
128+
def open_dataset(group: str) -> xr.Dataset:
129+
zarr_store = open_zarr_store(group)
130+
return xr.open_zarr(zarr_store, zarr_format=3, consolidated=False)
131+
132+
133+
def get_tile_center_coordiante(geobox: GeoBox, chunk: Chunk2D) -> tuple[float, float]:
134+
tile = geobox[chunk.y_start : chunk.y_end, chunk.x_start : chunk.x_end]
135+
center_coord = tile.to_crs("EPSG:4326").center_pixel.coordinates
136+
lat = center_coord["latitude"].values[0].item()
137+
lon = center_coord["longitude"].values[0].item()
138+
return lat, lon
139+
140+
141+
def normalize_latlon(lat: float, lon: float) -> tuple[tuple[float, float], tuple[float, float]]:
142+
lat = lat * np.pi / 180
143+
lon = lon * np.pi / 180
144+
145+
return (math.sin(lat), math.cos(lat)), (math.sin(lon), math.cos(lon))
146+
147+
148+
def normalize_timestamp(date: datetime) -> tuple[tuple[float, float], tuple[float, float]]:
149+
week = date.isocalendar().week * 2 * np.pi / 52
150+
hour = date.hour * 2 * np.pi / 24
151+
152+
return (math.sin(week), math.cos(week)), (math.sin(hour), math.cos(hour))
153+
154+
155+
def load_transform(bands: list[str], platform: str) -> tuple[Transform, list[float]]:
156+
with _CLAY_METADATA.open("r") as f:
157+
metadata = Box(yaml.safe_load(f))[platform]
158+
159+
mean = [metadata.bands.mean[band] for band in bands]
160+
std = [metadata.bands.std[band] for band in bands]
161+
wavelength = [metadata.bands.wavelength[band] for band in bands]
162+
163+
return Normalize(mean, std), wavelength
164+
165+
166+
class ClayInferenceTile(Task):
167+
chunk: Chunk2D
168+
mosaic_zarr_group: str
169+
roi: RegionOfInterest
170+
crs: str
171+
resolution: float
172+
output_zarr: tuple[str, str]
173+
174+
def execute(self, context: ExecutionContext) -> None:
175+
tracer = context._runner.tracer._tracer # type: ignore[arg-defined], # noqa: SLF001
176+
context.current_task.display = f"ClayInferenceTile({self.chunk})" # type: ignore[attr-defined]
177+
178+
with tracer.start_span("load_data"):
179+
start, end = self.roi.time
180+
start = datetime.fromisoformat(start)
181+
end = datetime.fromisoformat(end)
182+
mean_time = start + (end - start) / 2
183+
# the model takes the time of day into account, since we have a mosaic of lots of images we set it to noon
184+
# as an approximation of the middle of the day
185+
mean_time = mean_time.replace(hour=12, minute=0)
186+
lat, lon = get_tile_center_coordiante(self.roi.area.as_geobox(self.crs, self.resolution), self.chunk)
187+
188+
week_norm, hour_norm = normalize_timestamp(mean_time)
189+
lat_norm, lon_norm = normalize_latlon(lat, lon)
190+
191+
logger.info(f"Inference for tile {self.chunk} at lat={lat:.4f}, lon={lon:.4f} on {mean_time.isoformat()}")
192+
193+
cube = open_dataset(self.mosaic_zarr_group)
194+
bands = [s.item().decode("utf-8") for s in cube.band]
195+
transform, wavelengths = load_transform(bands, _CLAY_PLATFORM)
196+
197+
mosaic = cube.mosaic.isel(
198+
y=slice(self.chunk.y_start, self.chunk.y_end), x=slice(self.chunk.x_start, self.chunk.x_end)
199+
)
200+
201+
data = mosaic.load().to_numpy()
202+
# add a batch size
203+
data = np.expand_dims(data, axis=0)
204+
# convert to a contiguous array in float32
205+
data = np.ascontiguousarray(data.astype(np.float32))
206+
pixels = transform(torch.from_numpy(data))
207+
logger.info("Successfully loaded pixels")
208+
209+
model_input = {
210+
"platform": _CLAY_PLATFORM,
211+
"time": torch.tensor(
212+
np.hstack((week_norm, hour_norm)).reshape(1, 4),
213+
dtype=torch.float32,
214+
device=device(),
215+
),
216+
"latlon": torch.tensor(
217+
np.hstack((lat_norm, lon_norm)).reshape(1, 4), dtype=torch.float32, device=device()
218+
),
219+
"pixels": pixels.to(device()),
220+
"gsd": torch.tensor([self.resolution], device=device()),
221+
"waves": torch.tensor(wavelengths, device=device()),
222+
}
223+
224+
with tracer.start_span("load_model"):
225+
model = clay_model()
226+
227+
with tracer.start_span("inference"), torch.no_grad():
228+
unmsk_patch, _, _, _ = model.model.encoder(model_input)
229+
patches = unmsk_patch.detach().cpu().numpy()[0, 1:, :]
230+
patches = patches.reshape( # 32, 32, 1024
231+
CLAY_INFERENCE_TILE_SIZE // CLAY_PATCH_SIZE,
232+
CLAY_INFERENCE_TILE_SIZE // CLAY_PATCH_SIZE,
233+
CLAY_EMBEDDING_DIM,
234+
)
235+
236+
with tracer.start_span("write_output"):
237+
zarr_group_name, zarr_array_name = self.output_zarr
238+
zarr_group = zarr.open_group(open_zarr_store(zarr_group_name), mode="a")
239+
zarr_array: zarr.Array = zarr_group[zarr_array_name] # type: ignore[arg-type]
240+
zarr_array[
241+
self.chunk.y_start // CLAY_PATCH_SIZE : self.chunk.y_end // CLAY_PATCH_SIZE,
242+
self.chunk.x_start // CLAY_PATCH_SIZE : self.chunk.x_end // CLAY_PATCH_SIZE,
243+
:,
244+
] = patches
245+
246+
logger.info(f"Successfully wrote patches to Zarr array for tile {self.chunk}")
247+
248+
context.progress("inference").done(1)
249+
250+
251+
app = App()
252+
253+
254+
@app.default
255+
def main(cluster: str | None = None) -> None:
256+
if Path(".env").exists():
257+
assert load_dotenv()
258+
259+
service_name = f"{os.environ['RUNNER_NAME']}-{os.getpid()}"
260+
configure_console_logging()
261+
if os.environ.get("AXIOM_API_KEY"):
262+
configure_otel_logging_axiom(service_name)
263+
configure_otel_tracing_axiom(service_name)
264+
265+
client = WorkflowsClient() # a workflow client for https://api.tilebox.com
266+
267+
cache_client = boto3.client(
268+
"s3",
269+
endpoint_url="https://obs.eu-nl.otc.t-systems.com",
270+
aws_access_key_id=os.environ["OTC_ACCESS_KEY_ID"],
271+
aws_secret_access_key=os.environ["OTC_SECRET_ACCESS_KEY"],
272+
region_name="eu-nl",
273+
# without this boto will append x-amz-checksum-crc32:... to the contents of uploaded blobs
274+
config=Config(request_checksum_calculation="when_required", response_checksum_validation="when_required"),
275+
)
276+
cache = OTCBucketCache(OUTPUT_BUCKET, cache_client, prefix="cache/jobs")
277+
278+
logger.info(f"Starting runner on {cluster or 'default'} cluster")
279+
280+
runner = client.runner(
281+
cluster,
282+
tasks=[ClayInferenceOnMosaic, ClayInferenceTile],
283+
cache=cache,
284+
)
285+
runner.run_forever()
286+
287+
288+
if __name__ == "__main__":
289+
app()

‎s2-clay/infrastructure/Pulumi.dev.yaml‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,3 +5,15 @@ config:
55
secure: AAABADi4AtGJVRErlbS1Ff5IeLGtLWWzN0LUrH/Er/JWMYDbmMUVkYCDiT+So0q4PkY2uw==
66
tailscale:authKey:
77
secure: AAABAPnBHiuTwr/SBbMn3kEARr0B2+tt/4sJKWvUzCuofntyLlUgYaRdDjEmaHjWe7yDsNy8wzz/LRT25CENMuKzueBAfXUzyDkrF1mPIx/lrdMJjczx22FpuNH9
8+
runner:tileboxApiKey:
9+
secure: AAABABWf/0WqlIjt6L0qYZYl6sTLXvqi1LfFYqJxE5Y+JjRD5PEBqGe9mOECAdE+vtDDXKMV/jbBL1Bz9kA0dgVLRJjZo9PUxHqM5k+D5cfaLZd85ucaag0=
10+
runner:axiomApiKey:
11+
secure: AAABAKbcXSUX4sG5MrbhoDBT834Qg1Qmfbr58+951vac8bHagQ2SZDptYztRUAJi41/K2veF7X+vVvPr8l/8MsWAemB59BjMdw==
12+
runner:otcAccessKeyId:
13+
secure: AAABAJ/+jPAFit0crMZ1G+bw0bDBRtsdABbI2U6I7xh/UwhhJKIszxLEF3+ykZleiLjP3Q==
14+
runner:otcSecretAccessKey:
15+
secure: AAABANXIIyxfGQMoRfOvVAcFboRGl07LfC/ccefAeM9B1xpbh72urgcdbvsZ69QCqXHEWgzbWVJusuirZO3Oc6YR+cph6kWu
16+
runner:copernicusAccessKeyId:
17+
secure: AAABAJbfhOCq4dV7hqCt/MZxDj5s9TnY9m0oQKPPRGxhk97JQdGvhzC9Zux7+q1dM8K44w==
18+
runner:copernicusSecresAccessKey:
19+
secure: AAABAGOe66VNDwO7Bt6aJjF+ni35DcbSrXsXJ//7oOnUKRW9/SBf25NmvFMT5voMoE/k7Ywnt63hUrc5S+ZdLruJBRka7g0n

0 commit comments

Comments
 (0)