Spaces:
Running on Zero
Running on Zero
| import os | |
| import random | |
| import json | |
| import numpy as np | |
| from pyfastnoiselite.pyfastnoiselite import FastNoiseLite, NoiseType, FractalType | |
| import torch | |
| from terrain_diffusion.inference.perlin_transform import build_quantiles, transform_perlin | |
| STATS_CACHE_PATH = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(__file__))), "data", "global", "synthetic_map_stats.json") | |
| def _load_stats_cache(): | |
| if not os.path.exists(STATS_CACHE_PATH): | |
| return None | |
| try: | |
| with open(STATS_CACHE_PATH, "r", encoding="utf-8") as cache_file: | |
| data = json.load(cache_file) | |
| noise_quantile_tables = data["noise_quantile_tables"] | |
| data_quantile_tables = data["data_quantile_tables"] | |
| stats = { | |
| "a_temp_std": float(data["a_temp_std"]), | |
| "b_temp_std": float(data["b_temp_std"]), | |
| "temp_std_p1": float(data["temp_std_p1"]), | |
| "temp_std_p99": float(data["temp_std_p99"]), | |
| } | |
| for index, quantile_table in enumerate(noise_quantile_tables): | |
| stats[f"noise_quantiles_{index}"] = np.asarray(quantile_table, dtype=np.float64) | |
| for index, quantile_table in enumerate(data_quantile_tables): | |
| stats[f"base_image_quantiles_{index}"] = np.asarray(quantile_table, dtype=np.float64) | |
| print("Synthetic map stats cache hit.") | |
| return stats | |
| except Exception: | |
| print("Synthetic map stats cache unreadable. Recomputing.") | |
| return None | |
| def make_synthetic_map_factory(frequency_mult=[1.0, 1.0, 1.0, 1.0, 1.0], seed=None, drop_water_pct=0.0): | |
| actual_seeds = [((seed or random.randint(0, 2**30)) + i + 1) & 0x7FFFFFFF for i in range(5)] | |
| stats = _load_stats_cache() | |
| if stats is None: | |
| raise RuntimeError( | |
| f"Synthetic map stats not found at {STATS_CACHE_PATH}. " | |
| "This file must be pre-computed and included in the Space." | |
| ) | |
| a_temp_std = float(stats['a_temp_std']) | |
| b_temp_std = float(stats['b_temp_std']) | |
| temp_std_p1 = float(stats['temp_std_p1']) | |
| temp_std_p99 = float(stats['temp_std_p99']) | |
| def build_synthetic_map(frequency, octaves, lacunarity, gain, seed, noise_quantiles, base_image_quantiles): | |
| noise = FastNoiseLite(seed=seed) | |
| noise.noise_type = NoiseType.NoiseType_Perlin | |
| noise.frequency = frequency | |
| noise.fractal_type = FractalType.FractalType_FBm | |
| noise.fractal_octaves = octaves | |
| noise.fractal_lacunarity = lacunarity | |
| noise.fractal_gain = gain | |
| transform_fn = lambda x: transform_perlin(x, noise_quantiles, base_image_quantiles) | |
| return noise, transform_fn | |
| def sample_synthetic_map(noise, transform_fn, i1, j1, i2, j2): | |
| x = np.arange(i1, i2, dtype=np.float32) | |
| y = np.arange(j1, j2, dtype=np.float32) | |
| xx, yy = np.meshgrid(x, y) | |
| Xs = xx.flatten() | |
| Ys = yy.flatten() | |
| coords = np.array([Xs, Ys], dtype=np.float32) | |
| noise_values = noise.gen_from_coords(coords) | |
| transformed_values = transform_fn(noise_values) | |
| return transformed_values.reshape(i2 - i1, j2 - j1) | |
| map_configs = [ | |
| (0.05 * frequency_mult[0], 4, 2.0, 0.5), | |
| (0.05 * frequency_mult[1], 2, 2.0, 0.5), | |
| (0.05 * frequency_mult[2], 4, 2.0, 0.5), | |
| (0.05 * frequency_mult[3], 4, 2.0, 0.5), | |
| (0.05 * frequency_mult[4], 4, 2.0, 0.5), | |
| ] | |
| synthetic_params = [ | |
| build_synthetic_map(*cfg, actual_seeds[i], stats[f'noise_quantiles_{i}'], stats[f'base_image_quantiles_{i}']) | |
| for i, cfg in enumerate(map_configs) | |
| ] | |
| synthetic_elev_params, synthetic_temp_params, synthetic_temp_std_params, synthetic_precip_params, synthetic_precip_std_params = synthetic_params | |
| def finalize_synthetic_map(raw_map): | |
| synthetic_elev = np.asarray(raw_map[0], dtype=np.float32) | |
| synthetic_temp = np.asarray(raw_map[1], dtype=np.float32) | |
| synthetic_temp_std = np.asarray(raw_map[2], dtype=np.float32) | |
| synthetic_precip = np.asarray(raw_map[3], dtype=np.float32) | |
| synthetic_precip_std = np.asarray(raw_map[4], dtype=np.float32) | |
| lapse_rate = (-6.5 + 0.0015 * synthetic_precip).clip(-9.8, -4.0) / 1000 | |
| synthetic_temp = synthetic_temp + lapse_rate * np.maximum(0, synthetic_elev) | |
| synthetic_temp = np.clip(synthetic_temp, -10, 40) | |
| synthetic_temp = np.where(synthetic_temp > 20, synthetic_temp, (synthetic_temp - 20) * 1.25 + 20) | |
| t = (synthetic_temp_std - temp_std_p1) / (temp_std_p99 - temp_std_p1) | |
| baseline = np.maximum(temp_std_p1, -(a_temp_std * synthetic_temp + b_temp_std)) | |
| synthetic_temp_std = t * (temp_std_p99 - baseline) + baseline | |
| synthetic_temp_std = synthetic_temp_std + (a_temp_std * synthetic_temp + b_temp_std) | |
| synthetic_temp_std = np.maximum(synthetic_temp_std, 20) | |
| synthetic_precip_std = synthetic_precip_std * np.maximum(0, (185 - 0.04111 * synthetic_precip) / 185) | |
| return np.stack([synthetic_elev, synthetic_temp, synthetic_temp_std, synthetic_precip, synthetic_precip_std], axis=0) | |
| def sample_raw_synthetic_map(i1, j1, i2, j2): | |
| synthetic_elev = sample_synthetic_map(*synthetic_elev_params, i1, j1, i2, j2) | |
| synthetic_temp = sample_synthetic_map(*synthetic_temp_params, i1, j1, i2, j2) | |
| synthetic_temp_std = sample_synthetic_map(*synthetic_temp_std_params, i1, j1, i2, j2) | |
| synthetic_precip = sample_synthetic_map(*synthetic_precip_params, i1, j1, i2, j2) | |
| synthetic_precip_std = sample_synthetic_map(*synthetic_precip_std_params, i1, j1, i2, j2) | |
| return np.stack([synthetic_elev, synthetic_temp, synthetic_temp_std, synthetic_precip, synthetic_precip_std], axis=0) | |
| def sample_full_synthetic_map(i1, j1, i2, j2): | |
| synthetic_map = finalize_synthetic_map(sample_raw_synthetic_map(i1, j1, i2, j2)) | |
| synthetic_map[0] = np.sign(synthetic_map[0]) * np.sqrt(np.abs(synthetic_map[0])) | |
| return torch.from_numpy(synthetic_map).float() | |
| sample_full_synthetic_map.sample_raw = sample_raw_synthetic_map | |
| sample_full_synthetic_map.finalize = finalize_synthetic_map | |
| return sample_full_synthetic_map | |