Download scripts/fake_data.py from OneScience-Group/AlphaEarthFoundations: direct link, hf CLI and curl.
- Browser
- Download file 5.85 kB
-
https://huggingface.co/OneScience-Group/AlphaEarthFoundations/resolve/main/scripts/fake_data.py
- Command line
-
hf download hf://OneScience-Group/AlphaEarthFoundations/scripts/fake_data.py
-
curl -L -o fake_data.py https://huggingface.co/OneScience-Group/AlphaEarthFoundations/resolve/main/scripts/fake_data.py
5.85 kB
| """Generate sparse synthetic observations following the paper-confirmed data contract.""" | |
| import argparse | |
| from pathlib import Path | |
| import numpy as np | |
| import yaml | |
| ROOT = Path(__file__).resolve().parents[1] | |
| DAY_MS = 86_400_000 | |
| TARGET_ONLY_SLOTS = 3 | |
| def load_config(): | |
| return yaml.safe_load((ROOT / "conf/config.yaml").read_text()) | |
| def standardized_signal(rng, count, frames, channels, size, source_phase): | |
| y, x = np.mgrid[-1:1:complex(size), -1:1:complex(size)].astype(np.float32) | |
| base = np.sin(3 * np.pi * x) * np.cos(2 * np.pi * y) + 0.4 * x + 0.2 * y | |
| values = np.empty((count, frames, channels, size, size), np.float32) | |
| for sample in range(count): | |
| phase = rng.uniform(0, 2 * np.pi) + source_phase | |
| for step in range(frames): | |
| seasonal = np.sin(2 * np.pi * step / max(frames, 1) + phase) | |
| for channel in range(channels): | |
| values[sample, step, channel] = base + 0.08 * channel + 0.25 * seasonal | |
| values[sample] += rng.normal(0, 0.04, values[sample].shape) | |
| values -= values.mean(axis=(0, 1, 3, 4), keepdims=True) | |
| values /= values.std(axis=(0, 1, 3, 4), keepdims=True).clip(1e-6) | |
| return np.clip(values, -6, 6).astype(np.float32) | |
| def quality_mask(rng, count, frames, channels, size, sparse=False): | |
| probability = 0.08 if not sparse else 0.92 | |
| mask = rng.random((count, frames, 1, size, size)) > probability | |
| if not sparse: | |
| mask[:, :, :, :2] = False | |
| return np.repeat(mask, channels, axis=2).astype(np.float32) | |
| def make_split(path, count, config, seed): | |
| rng = np.random.default_rng(seed) | |
| data, size = config["data"], config["data"]["image_size"] | |
| start = np.datetime64("2020-01-01", "ms").astype(np.int64) | |
| payload = {"format_version": np.asarray(data["format_version"]), "data_source": np.asarray("synthetic")} | |
| for source_index, (name, spec) in enumerate(data["input_sources"].items()): | |
| frames, channels = spec["timesteps"], spec["channels"] | |
| payload[name] = standardized_signal(rng, count, frames, channels, size, source_index) | |
| days = np.linspace(0, 364, frames, dtype=np.int64) | |
| payload[f"timestamps_{name}"] = np.tile(start + days * DAY_MS + source_index, (count, 1)) | |
| available = np.ones((count, frames), np.bool_) | |
| available[:, -max(1, frames // 10):] = False | |
| payload[f"frame_available_{name}"] = available | |
| channel_available = np.ones((count, frames, channels), np.bool_) | |
| if name == "sentinel1": | |
| channel_available[..., :4] = False | |
| for sample in range(count): | |
| for step in range(frames): | |
| pair = (0, 1) if (sample + step) % 2 == 0 else (2, 3) | |
| channel_available[sample, step, list(pair)] = True | |
| payload[f"channel_available_{name}"] = channel_available | |
| payload[f"pixel_valid_{name}"] = quality_mask(rng, count, frames, channels, size) | |
| geometry = np.zeros((count, frames, config["model"]["max_geometry_dim"]), np.float32) | |
| if name == "sentinel1": | |
| geometry[..., 0] = rng.uniform(-np.pi, np.pi, (count, frames)) | |
| geometry[..., 1] = rng.uniform(1.6, 1.8, (count, frames)) | |
| payload[f"geometry_{name}"] = geometry | |
| for target_index, (name, spec) in enumerate(data["target_sources"].items()): | |
| if name in data["input_sources"]: | |
| continue | |
| channels = 1 if spec["type"] == "categorical" else spec["channels"] | |
| values = standardized_signal(rng, count, TARGET_ONLY_SLOTS, channels, size, target_index + 3) | |
| if spec["type"] == "categorical": | |
| bins = np.linspace(values.min(), values.max(), spec["channels"] + 1)[1:-1] | |
| values = np.digitize(values[:, :, 0], bins).astype(np.int64) | |
| payload[f"target_sequence_{name}"] = values | |
| payload[f"target_timestamps_{name}"] = np.tile( | |
| start + np.asarray([30, 182, 330], np.int64) * DAY_MS + target_index, (count, 1) | |
| ) | |
| mask_channels = 1 if spec["type"] == "categorical" else spec["channels"] | |
| payload[f"target_pixel_valid_{name}"] = quality_mask( | |
| rng, count, TARGET_ONLY_SLOTS, mask_channels, size, sparse=name == "gedi" | |
| ) | |
| geometry = np.zeros((count, TARGET_ONLY_SLOTS, config["model"]["max_geometry_dim"]), np.float32) | |
| if name == "palsar2": | |
| geometry[..., 0] = rng.integers(0, 2, (count, TARGET_ONLY_SLOTS)) | |
| geometry[..., 1] = rng.integers(0, 2, (count, TARGET_ONLY_SLOTS)) | |
| payload[f"target_geometry_{name}"] = geometry | |
| payload[f"target_frame_available_{name}"] = np.ones((count, TARGET_ONLY_SLOTS), np.bool_) | |
| payload["support_period"] = np.tile(np.array([start, start + 365 * DAY_MS], np.int64), (count, 1)) | |
| payload["valid_period"] = np.tile(np.array([start + 60 * DAY_MS, start + 300 * DAY_MS], np.int64), (count, 1)) | |
| text = rng.normal(size=(count, config["model"]["embedding_dim"])).astype(np.float32) | |
| payload["text_target"] = text / np.linalg.norm(text, axis=1, keepdims=True) | |
| np.savez_compressed(path, **payload) | |
| def main(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--force", action="store_true") | |
| args = parser.parse_args() | |
| config = load_config() | |
| data_dir = ROOT / config["data"]["root"] | |
| data_dir.mkdir(parents=True, exist_ok=True) | |
| for offset, (name, count) in enumerate((("train.npz", config["data"]["train_samples"]), | |
| ("test.npz", config["data"]["test_samples"]))): | |
| target = data_dir / name | |
| if args.force or not target.exists(): | |
| make_split(target, count, config, config["seed"] + offset) | |
| print(f"generated={target.relative_to(ROOT)} samples={count} format={config['data']['format_version']}") | |
| if __name__ == "__main__": | |
| main() | |