/
sususer
/
ColonyGEN
Обзор
Документация
Войти
/
sususer
/
ColonyGEN
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
main
ml/vae/model.py
696 строк
31 KB
Chekr
f11
03 июн 2026, 13:53
03 июн 2026, 13:53
d260831
Код
Авторство
О чём код?
from __future__ import annotations import json from dataclasses import dataclass from pathlib import Path from typing import Any import numpy as np @dataclass class VAEConfig: map_width: int map_height: int tile_feature_dim: int condition_dim: int = 0 terrain_feature_dim: int = 7 biome_feature_dim: int = 11 layer_feature_dim: int = 5 tile_embedding_dim: int = 6 encoder_hidden_dim: int = 256 bottleneck_hidden_dim: int = 96 latent_dim: int = 32 seed: int = 42 def __post_init__(self) -> None: expected = self.terrain_feature_dim + self.biome_feature_dim + 2 + self.layer_feature_dim if expected != self.tile_feature_dim: raise ValueError( f"Invalid VAEConfig feature layout: expected {expected}, got {self.tile_feature_dim}" ) @property def num_tiles(self) -> int: return self.map_width * self.map_height @property def embedding_dim(self) -> int: return self.num_tiles * self.tile_embedding_dim @property def input_dim(self) -> int: return self.num_tiles * self.tile_feature_dim @property def terrain_slice(self) -> slice: return slice(0, self.terrain_feature_dim) @property def biome_slice(self) -> slice: return slice(self.terrain_slice.stop, self.terrain_slice.stop + self.biome_feature_dim) @property def resource_index(self) -> int: return self.biome_slice.stop @property def base_index(self) -> int: return self.resource_index + 1 @property def numeric_slice(self) -> slice: return slice(self.resource_index, self.tile_feature_dim) def to_dict(self) -> dict[str, Any]: return { "map_width": self.map_width, "map_height": self.map_height, "tile_feature_dim": self.tile_feature_dim, "condition_dim": self.condition_dim, "terrain_feature_dim": self.terrain_feature_dim, "biome_feature_dim": self.biome_feature_dim, "layer_feature_dim": self.layer_feature_dim, "tile_embedding_dim": self.tile_embedding_dim, "encoder_hidden_dim": self.encoder_hidden_dim, "bottleneck_hidden_dim": self.bottleneck_hidden_dim, "latent_dim": self.latent_dim, "seed": self.seed, } @classmethod def from_dict(cls, data: dict[str, Any]) -> "VAEConfig": return cls( map_width=int(data["map_width"]), map_height=int(data["map_height"]), tile_feature_dim=int(data["tile_feature_dim"]), condition_dim=int(data.get("condition_dim", 0)), terrain_feature_dim=int(data.get("terrain_feature_dim", 7)), biome_feature_dim=int(data.get("biome_feature_dim", 11)), layer_feature_dim=int(data.get("layer_feature_dim", 5)), tile_embedding_dim=int(data.get("tile_embedding_dim", 6)), encoder_hidden_dim=int(data.get("encoder_hidden_dim", 256)), bottleneck_hidden_dim=int(data.get("bottleneck_hidden_dim", 96)), latent_dim=int(data.get("latent_dim", 32)), seed=int(data.get("seed", 42)), ) class AdamOptimizer: def __init__( self, learning_rate: float = 1e-3, beta1: float = 0.9, beta2: float = 0.999, epsilon: float = 1e-8, weight_decay: float = 0.0, clip_norm: float | None = 5.0, ): self.learning_rate = learning_rate self.beta1 = beta1 self.beta2 = beta2 self.epsilon = epsilon self.weight_decay = weight_decay self.clip_norm = clip_norm self.step_count = 0 self.m: dict[str, np.ndarray] = {} self.v: dict[str, np.ndarray] = {} def step(self, params: dict[str, np.ndarray], grads: dict[str, np.ndarray]) -> float: self.step_count += 1 global_norm = float(np.sqrt(sum(np.sum(grad * grad) for grad in grads.values()))) scale = 1.0 if self.clip_norm is not None and global_norm > self.clip_norm: scale = self.clip_norm / (global_norm + 1e-12) for name, param in params.items(): grad = grads[name].astype(np.float32) * scale if self.weight_decay: grad = grad + self.weight_decay * param if name not in self.m: self.m[name] = np.zeros_like(param) self.v[name] = np.zeros_like(param) self.m[name] = self.beta1 * self.m[name] + (1.0 - self.beta1) * grad self.v[name] = self.beta2 * self.v[name] + (1.0 - self.beta2) * (grad * grad) m_hat = self.m[name] / (1.0 - self.beta1 ** self.step_count) v_hat = self.v[name] / (1.0 - self.beta2 ** self.step_count) params[name] -= self.learning_rate * m_hat / (np.sqrt(v_hat) + self.epsilon) return global_norm def relu(x: np.ndarray) -> np.ndarray: return np.maximum(x, 0.0) def relu_grad(x: np.ndarray) -> np.ndarray: return (x > 0.0).astype(np.float32) def sigmoid(x: np.ndarray) -> np.ndarray: clipped = np.clip(x, -30.0, 30.0) return 1.0 / (1.0 + np.exp(-clipped)) def softmax(x: np.ndarray, axis: int = -1) -> np.ndarray: shifted = x - np.max(x, axis=axis, keepdims=True) exp = np.exp(np.clip(shifted, -30.0, 30.0)) return exp / np.sum(exp, axis=axis, keepdims=True) def terrain_resource_mask(terrain_output: np.ndarray, terrain_slice: slice) -> np.ndarray: forest_idx = 2 - terrain_slice.start mountain_idx = 3 - terrain_slice.start forest_prob = terrain_output[:, :, forest_idx:forest_idx + 1] mountain_prob = terrain_output[:, :, mountain_idx:mountain_idx + 1] return np.clip(forest_prob + mountain_prob, 0.0, 1.0) def neighbor_sum4(values: np.ndarray, height: int, width: int) -> np.ndarray: reshaped = values.reshape(values.shape[0], height, width, values.shape[-1]) padded = np.pad(reshaped, ((0, 0), (1, 1), (1, 1), (0, 0)), mode="constant") total = ( padded[:, :-2, 1:-1, :] + padded[:, 2:, 1:-1, :] + padded[:, 1:-1, :-2, :] + padded[:, 1:-1, 2:, :] ) return total.reshape(values.shape[0], height * width, values.shape[-1]) def forest_dispersion_mask( forest_prob: np.ndarray, resource_signal: np.ndarray, height: int, width: int, ) -> np.ndarray: forest_drive = forest_prob * resource_signal crowding = neighbor_sum4(forest_drive, height, width) / 4.0 canopy_support = neighbor_sum4(forest_prob, height, width) / 4.0 dispersion = 0.92 - crowding * 0.55 + canopy_support * 0.22 return np.clip(dispersion, 0.40, 1.0) def hydrology_water_bias(numeric_output_raw: np.ndarray, config: VAEConfig) -> np.ndarray: layer_offset = config.base_index + 1 - config.numeric_slice.start elevation = numeric_output_raw[:, :, layer_offset:layer_offset + 1] moisture = numeric_output_raw[:, :, layer_offset + 1:layer_offset + 2] lowland_signal = np.clip((0.60 - elevation) / 0.32, 0.0, 1.0) wet_signal = np.clip((moisture - 0.40) / 0.28, 0.0, 1.0) basin_signal = np.clip((0.50 - elevation) / 0.20 + (moisture - 0.46) / 0.24, 0.0, 1.0) return np.clip(lowland_signal * 0.45 + wet_signal * 0.30 + basin_signal * 0.40, 0.0, 1.0) class NumpyVAE: def __init__( self, config: VAEConfig, feature_weights: np.ndarray | None = None, terrain_class_weights: np.ndarray | None = None, ): self.config = config self.rng = np.random.default_rng(config.seed) self.params: dict[str, np.ndarray] = {} self.feature_weights = np.asarray( feature_weights if feature_weights is not None else np.ones(config.tile_feature_dim, dtype=np.float32), dtype=np.float32, ) self.terrain_class_weights = np.asarray( terrain_class_weights if terrain_class_weights is not None else np.ones(config.terrain_feature_dim, dtype=np.float32), dtype=np.float32, ) self._init_parameters() def _init_parameters(self) -> None: c = self.config self.params["W_tile_in"] = self._init_weight(c.tile_feature_dim, c.tile_embedding_dim) self.params["b_tile_in"] = np.zeros((c.tile_embedding_dim,), dtype=np.float32) self.params["W_enc1"] = self._init_weight(c.embedding_dim, c.encoder_hidden_dim) if c.condition_dim > 0: self.params["W_cond_enc1"] = self._init_weight(c.condition_dim, c.encoder_hidden_dim) self.params["b_enc1"] = np.zeros((c.encoder_hidden_dim,), dtype=np.float32) self.params["W_enc2"] = self._init_weight(c.encoder_hidden_dim, c.bottleneck_hidden_dim) if c.condition_dim > 0: self.params["W_cond_enc2"] = self._init_weight(c.condition_dim, c.bottleneck_hidden_dim) self.params["b_enc2"] = np.zeros((c.bottleneck_hidden_dim,), dtype=np.float32) self.params["W_mu"] = self._init_weight(c.bottleneck_hidden_dim, c.latent_dim) if c.condition_dim > 0: self.params["W_cond_mu"] = self._init_weight(c.condition_dim, c.latent_dim) self.params["b_mu"] = np.zeros((c.latent_dim,), dtype=np.float32) self.params["W_logvar"] = self._init_weight(c.bottleneck_hidden_dim, c.latent_dim) if c.condition_dim > 0: self.params["W_cond_logvar"] = self._init_weight(c.condition_dim, c.latent_dim) self.params["b_logvar"] = np.zeros((c.latent_dim,), dtype=np.float32) if c.condition_dim > 0: self.params["W_prior_mu"] = self._init_weight(c.condition_dim, c.latent_dim) self.params["b_prior_mu"] = np.zeros((c.latent_dim,), dtype=np.float32) self.params["W_prior_logvar"] = self._init_weight(c.condition_dim, c.latent_dim) self.params["b_prior_logvar"] = np.zeros((c.latent_dim,), dtype=np.float32) self.params["W_dec1"] = self._init_weight(c.latent_dim, c.bottleneck_hidden_dim) if c.condition_dim > 0: self.params["W_cond_dec1"] = self._init_weight(c.condition_dim, c.bottleneck_hidden_dim) self.params["b_dec1"] = np.zeros((c.bottleneck_hidden_dim,), dtype=np.float32) self.params["W_dec2"] = self._init_weight(c.bottleneck_hidden_dim, c.encoder_hidden_dim) if c.condition_dim > 0: self.params["W_cond_dec2"] = self._init_weight(c.condition_dim, c.encoder_hidden_dim) self.params["b_dec2"] = np.zeros((c.encoder_hidden_dim,), dtype=np.float32) self.params["W_dec3"] = self._init_weight(c.encoder_hidden_dim, c.embedding_dim) self.params["b_dec3"] = np.zeros((c.embedding_dim,), dtype=np.float32) self.params["W_tile_out"] = self._init_weight(c.tile_embedding_dim, c.tile_feature_dim) if c.condition_dim > 0: self.params["W_cond_out"] = self._init_weight(c.condition_dim, c.tile_feature_dim) self.params["b_tile_out"] = np.zeros((c.tile_feature_dim,), dtype=np.float32) def _init_weight(self, fan_in: int, fan_out: int) -> np.ndarray: limit = np.sqrt(6.0 / float(fan_in + fan_out)) return self.rng.uniform(-limit, limit, size=(fan_in, fan_out)).astype(np.float32) def _decode_output_from_logits(self, out_pre: np.ndarray, training: bool) -> dict[str, np.ndarray]: c = self.config output = np.empty_like(out_pre) numeric_output_raw = sigmoid(out_pre[:, :, c.numeric_slice]) terrain_logits = out_pre[:, :, c.terrain_slice].copy() if not training: water_idx = 6 - c.terrain_slice.start plain_idx = 4 - c.terrain_slice.start sand_idx = 5 - c.terrain_slice.start water_bias = hydrology_water_bias(numeric_output_raw, c) terrain_logits[:, :, water_idx:water_idx + 1] += water_bias * 1.35 terrain_logits[:, :, plain_idx:plain_idx + 1] -= water_bias * 0.24 terrain_logits[:, :, sand_idx:sand_idx + 1] -= water_bias * 0.16 terrain_output = softmax(terrain_logits, axis=-1) output[:, :, c.terrain_slice] = terrain_output output[:, :, c.biome_slice] = softmax(out_pre[:, :, c.biome_slice], axis=-1) forest_idx = 2 - c.terrain_slice.start mountain_idx = 3 - c.terrain_slice.start forest_prob = terrain_output[:, :, forest_idx:forest_idx + 1] mountain_prob = terrain_output[:, :, mountain_idx:mountain_idx + 1] resource_offset = c.resource_index - c.numeric_slice.start resource_signal = numeric_output_raw[:, :, resource_offset:resource_offset + 1] forest_mask = forest_prob * forest_dispersion_mask( forest_prob, resource_signal, c.map_height, c.map_width, ) resource_mask = np.clip(mountain_prob + forest_mask, 0.0, 1.0) numeric_output = numeric_output_raw.copy() numeric_output[:, :, resource_offset:resource_offset + 1] *= resource_mask output[:, :, c.numeric_slice] = numeric_output return { "output": output, "terrain_output": terrain_output, "numeric_output_raw": numeric_output_raw, "resource_mask": resource_mask, "forest_resource_mask": forest_mask, } def _prepare_condition(self, condition: np.ndarray | None, batch_size: int) -> np.ndarray: c = self.config if c.condition_dim <= 0: return np.zeros((batch_size, 0), dtype=np.float32) if condition is None: return np.zeros((batch_size, c.condition_dim), dtype=np.float32) array = np.asarray(condition, dtype=np.float32) if array.ndim == 1: array = np.repeat(array.reshape(1, -1), batch_size, axis=0) if array.shape != (batch_size, c.condition_dim): raise ValueError( f"Invalid condition shape: expected {(batch_size, c.condition_dim)}, got {array.shape}" ) return array def _conditional_prior(self, condition_array: np.ndarray) -> tuple[np.ndarray, np.ndarray]: c = self.config if c.condition_dim <= 0: batch_size = condition_array.shape[0] zeros = np.zeros((batch_size, c.latent_dim), dtype=np.float32) return zeros, zeros p = self.params prior_mu = condition_array @ p["W_prior_mu"] + p["b_prior_mu"] prior_logvar = condition_array @ p["W_prior_logvar"] + p["b_prior_logvar"] return prior_mu.astype(np.float32), prior_logvar.astype(np.float32) def _reconstruction_loss( self, x_grid: np.ndarray, out_pre: np.ndarray, output: np.ndarray, numeric_output_raw: np.ndarray, resource_mask: np.ndarray, with_grad: bool, ) -> tuple[float, np.ndarray | None]: c = self.config batch_size = x_grid.shape[0] denom = float(batch_size * c.num_tiles) eps = 1e-7 terrain_target = x_grid[:, :, c.terrain_slice] terrain_output = output[:, :, c.terrain_slice] terrain_weight = float(np.mean(self.feature_weights[c.terrain_slice])) terrain_class_weights = self.terrain_class_weights.reshape(1, 1, -1) terrain_loss = -terrain_weight * float( np.sum(terrain_class_weights * terrain_target * np.log(np.clip(terrain_output, eps, 1.0))) / denom ) biome_target = x_grid[:, :, c.biome_slice] biome_output = output[:, :, c.biome_slice] biome_weight = float(np.mean(self.feature_weights[c.biome_slice])) biome_loss = -biome_weight * float( np.sum(biome_target * np.log(np.clip(biome_output, eps, 1.0))) / denom ) numeric_target = x_grid[:, :, c.numeric_slice] numeric_output = output[:, :, c.numeric_slice] numeric_weights = self.feature_weights[c.numeric_slice].reshape(1, 1, -1) numeric_diff = numeric_output - numeric_target numeric_loss = float(np.sum(numeric_weights * numeric_diff * numeric_diff) / denom) grad_out_pre = None if with_grad: grad_out_pre = np.zeros_like(out_pre) grad_out_pre[:, :, c.terrain_slice] = ( terrain_weight * terrain_class_weights * (terrain_output - terrain_target) / denom ) grad_out_pre[:, :, c.biome_slice] = ( biome_weight * (biome_output - biome_target) / denom ) grad_numeric = (2.0 * numeric_weights * numeric_diff) / denom grad_numeric_pre = grad_numeric * numeric_output_raw * (1.0 - numeric_output_raw) resource_offset = c.resource_index - c.numeric_slice.start grad_numeric_pre[:, :, resource_offset:resource_offset + 1] *= resource_mask grad_out_pre[:, :, c.numeric_slice] = grad_numeric_pre return terrain_loss + biome_loss + numeric_loss, grad_out_pre def forward( self, x: np.ndarray, training: bool, condition: np.ndarray | None = None, seed: int | None = None, ) -> dict[str, np.ndarray]: c = self.config p = self.params batch_size = x.shape[0] condition_array = self._prepare_condition(condition, batch_size) x_grid = x.reshape(batch_size, c.num_tiles, c.tile_feature_dim) tile_pre = np.einsum("bnf,fe->bne", x_grid, p["W_tile_in"]) + p["b_tile_in"] tile_hidden = np.tanh(tile_pre) flat = tile_hidden.reshape(batch_size, c.embedding_dim) enc1_pre = flat @ p["W_enc1"] if c.condition_dim > 0: enc1_pre = enc1_pre + condition_array @ p["W_cond_enc1"] enc1_pre = enc1_pre + p["b_enc1"] enc1 = relu(enc1_pre) enc2_pre = enc1 @ p["W_enc2"] if c.condition_dim > 0: enc2_pre = enc2_pre + condition_array @ p["W_cond_enc2"] enc2_pre = enc2_pre + p["b_enc2"] enc2 = relu(enc2_pre) prior_mu, prior_logvar = self._conditional_prior(condition_array) mu = enc2 @ p["W_mu"] logvar = enc2 @ p["W_logvar"] if c.condition_dim > 0: mu = mu + condition_array @ p["W_cond_mu"] logvar = logvar + condition_array @ p["W_cond_logvar"] mu = mu + p["b_mu"] logvar = logvar + p["b_logvar"] std = np.exp(0.5 * np.clip(logvar, -20.0, 20.0)).astype(np.float32) if training: rng = np.random.default_rng(seed) eps = rng.normal(0.0, 1.0, size=mu.shape).astype(np.float32) z = mu + eps * std else: eps = np.zeros_like(mu, dtype=np.float32) z = mu dec1_pre = z @ p["W_dec1"] if c.condition_dim > 0: dec1_pre = dec1_pre + condition_array @ p["W_cond_dec1"] dec1_pre = dec1_pre + p["b_dec1"] dec1 = relu(dec1_pre) dec2_pre = dec1 @ p["W_dec2"] if c.condition_dim > 0: dec2_pre = dec2_pre + condition_array @ p["W_cond_dec2"] dec2_pre = dec2_pre + p["b_dec2"] dec2 = relu(dec2_pre) flat_dec = dec2 @ p["W_dec3"] + p["b_dec3"] tile_dec_pre = flat_dec.reshape(batch_size, c.num_tiles, c.tile_embedding_dim) tile_dec = np.tanh(tile_dec_pre) out_pre = np.einsum("bne,ef->bnf", tile_dec, p["W_tile_out"]) + p["b_tile_out"] if c.condition_dim > 0: out_pre = out_pre + (condition_array @ p["W_cond_out"]).reshape(batch_size, 1, c.tile_feature_dim) decoded = self._decode_output_from_logits(out_pre, training=training) output = decoded["output"] return { "x_grid": x_grid, "tile_pre": tile_pre, "tile_hidden": tile_hidden, "flat": flat, "condition": condition_array, "enc1_pre": enc1_pre, "enc1": enc1, "enc2_pre": enc2_pre, "enc2": enc2, "prior_mu": prior_mu, "prior_logvar": prior_logvar, "mu": mu, "logvar": logvar, "std": std, "eps": eps, "z": z, "dec1_pre": dec1_pre, "dec1": dec1, "dec2_pre": dec2_pre, "dec2": dec2, "flat_dec": flat_dec, "tile_dec_pre": tile_dec_pre, "tile_dec": tile_dec, "out_pre": out_pre, "output": output, "numeric_output_raw": decoded["numeric_output_raw"], "resource_mask": decoded["resource_mask"], "forest_resource_mask": decoded["forest_resource_mask"], } def loss_and_gradients( self, x: np.ndarray, condition: np.ndarray | None = None, beta: float = 1.0, seed: int | None = None, ) -> tuple[dict[str, float], dict[str, np.ndarray]]: c = self.config p = self.params cache = self.forward(x, training=True, condition=condition, seed=seed) batch_size = x.shape[0] x_grid = cache["x_grid"] recon_loss, grad_out_pre = self._reconstruction_loss( x_grid, cache["out_pre"], cache["output"], cache["numeric_output_raw"], cache["resource_mask"], with_grad=True, ) posterior_var = np.exp(np.clip(cache["logvar"], -20.0, 20.0)) if c.condition_dim > 0: prior_var = np.exp(np.clip(cache["prior_logvar"], -20.0, 20.0)) mu_delta = cache["mu"] - cache["prior_mu"] kl_loss = float( 0.5 * np.sum( cache["prior_logvar"] - cache["logvar"] + (posterior_var + mu_delta ** 2) / prior_var - 1.0 ) / batch_size ) else: prior_var = np.ones_like(posterior_var, dtype=np.float32) mu_delta = cache["mu"] kl_loss = float( 0.5 * np.sum(posterior_var + cache["mu"] ** 2 - 1.0 - cache["logvar"]) / batch_size ) total_loss = recon_loss + beta * kl_loss grads: dict[str, np.ndarray] = {} grads["W_tile_out"] = np.einsum("bne,bnf->ef", cache["tile_dec"], grad_out_pre).astype(np.float32) if c.condition_dim > 0: grads["W_cond_out"] = (cache["condition"].T @ np.sum(grad_out_pre, axis=1)).astype(np.float32) grads["b_tile_out"] = np.sum(grad_out_pre, axis=(0, 1)).astype(np.float32) grad_tile_dec = np.einsum("bnf,ef->bne", grad_out_pre, p["W_tile_out"]) grad_tile_dec_pre = grad_tile_dec * (1.0 - cache["tile_dec"] ** 2) grad_flat_dec = grad_tile_dec_pre.reshape(batch_size, c.embedding_dim) grads["W_dec3"] = (cache["dec2"].T @ grad_flat_dec).astype(np.float32) grads["b_dec3"] = np.sum(grad_flat_dec, axis=0).astype(np.float32) grad_dec2 = grad_flat_dec @ p["W_dec3"].T grad_dec2_pre = grad_dec2 * relu_grad(cache["dec2_pre"]) grads["W_dec2"] = (cache["dec1"].T @ grad_dec2_pre).astype(np.float32) if c.condition_dim > 0: grads["W_cond_dec2"] = (cache["condition"].T @ grad_dec2_pre).astype(np.float32) grads["b_dec2"] = np.sum(grad_dec2_pre, axis=0).astype(np.float32) grad_dec1 = grad_dec2_pre @ p["W_dec2"].T grad_dec1_pre = grad_dec1 * relu_grad(cache["dec1_pre"]) grads["W_dec1"] = (cache["z"].T @ grad_dec1_pre).astype(np.float32) if c.condition_dim > 0: grads["W_cond_dec1"] = (cache["condition"].T @ grad_dec1_pre).astype(np.float32) grads["b_dec1"] = np.sum(grad_dec1_pre, axis=0).astype(np.float32) grad_z = grad_dec1_pre @ p["W_dec1"].T grad_mu = grad_z + (beta * mu_delta / (prior_var * batch_size)) grad_logvar = grad_z * cache["eps"] * 0.5 * cache["std"] grad_logvar += 0.5 * beta * (posterior_var / prior_var - 1.0) / batch_size grads["W_mu"] = (cache["enc2"].T @ grad_mu).astype(np.float32) if c.condition_dim > 0: grads["W_cond_mu"] = (cache["condition"].T @ grad_mu).astype(np.float32) grads["b_mu"] = np.sum(grad_mu, axis=0).astype(np.float32) grads["W_logvar"] = (cache["enc2"].T @ grad_logvar).astype(np.float32) if c.condition_dim > 0: grads["W_cond_logvar"] = (cache["condition"].T @ grad_logvar).astype(np.float32) grads["b_logvar"] = np.sum(grad_logvar, axis=0).astype(np.float32) if c.condition_dim > 0: grad_prior_mu = beta * (cache["prior_mu"] - cache["mu"]) / (prior_var * batch_size) grad_prior_logvar = 0.5 * beta * (1.0 - (posterior_var + mu_delta ** 2) / prior_var) / batch_size grads["W_prior_mu"] = (cache["condition"].T @ grad_prior_mu).astype(np.float32) grads["b_prior_mu"] = np.sum(grad_prior_mu, axis=0).astype(np.float32) grads["W_prior_logvar"] = (cache["condition"].T @ grad_prior_logvar).astype(np.float32) grads["b_prior_logvar"] = np.sum(grad_prior_logvar, axis=0).astype(np.float32) grad_enc2 = grad_mu @ p["W_mu"].T + grad_logvar @ p["W_logvar"].T grad_enc2_pre = grad_enc2 * relu_grad(cache["enc2_pre"]) grads["W_enc2"] = (cache["enc1"].T @ grad_enc2_pre).astype(np.float32) if c.condition_dim > 0: grads["W_cond_enc2"] = (cache["condition"].T @ grad_enc2_pre).astype(np.float32) grads["b_enc2"] = np.sum(grad_enc2_pre, axis=0).astype(np.float32) grad_enc1 = grad_enc2_pre @ p["W_enc2"].T grad_enc1_pre = grad_enc1 * relu_grad(cache["enc1_pre"]) grads["W_enc1"] = (cache["flat"].T @ grad_enc1_pre).astype(np.float32) if c.condition_dim > 0: grads["W_cond_enc1"] = (cache["condition"].T @ grad_enc1_pre).astype(np.float32) grads["b_enc1"] = np.sum(grad_enc1_pre, axis=0).astype(np.float32) grad_flat = grad_enc1_pre @ p["W_enc1"].T grad_tile_hidden = grad_flat.reshape(batch_size, c.num_tiles, c.tile_embedding_dim) grad_tile_pre = grad_tile_hidden * (1.0 - cache["tile_hidden"] ** 2) grads["W_tile_in"] = np.einsum("bnf,bne->fe", x_grid, grad_tile_pre).astype(np.float32) grads["b_tile_in"] = np.sum(grad_tile_pre, axis=(0, 1)).astype(np.float32) metrics = { "loss": total_loss, "reconstruction_loss": recon_loss, "kl_loss": kl_loss, } return metrics, grads def evaluate_batch(self, x: np.ndarray, condition: np.ndarray | None = None, beta: float = 1.0) -> dict[str, float]: c = self.config cache = self.forward(x, training=False, condition=condition) batch_size = x.shape[0] x_grid = cache["x_grid"] recon_loss, _ = self._reconstruction_loss( x_grid, cache["out_pre"], cache["output"], cache["numeric_output_raw"], cache["resource_mask"], with_grad=False, ) posterior_var = np.exp(np.clip(cache["logvar"], -20.0, 20.0)) if c.condition_dim > 0: prior_var = np.exp(np.clip(cache["prior_logvar"], -20.0, 20.0)) mu_delta = cache["mu"] - cache["prior_mu"] kl_loss = float( 0.5 * np.sum( cache["prior_logvar"] - cache["logvar"] + (posterior_var + mu_delta ** 2) / prior_var - 1.0 ) / batch_size ) else: kl_loss = float( 0.5 * np.sum(posterior_var + cache["mu"] ** 2 - 1.0 - cache["logvar"]) / batch_size ) return { "loss": recon_loss + beta * kl_loss, "reconstruction_loss": recon_loss, "kl_loss": kl_loss, } def reconstruct(self, x: np.ndarray, condition: np.ndarray | None = None) -> np.ndarray: return self.forward(x, training=False, condition=condition)["output"].reshape(x.shape[0], -1) def encode(self, x: np.ndarray, condition: np.ndarray | None = None) -> tuple[np.ndarray, np.ndarray]: cache = self.forward(x, training=False, condition=condition) return cache["mu"], cache["logvar"] def decode(self, z: np.ndarray, condition: np.ndarray | None = None) -> np.ndarray: p = self.params c = self.config latent = np.asarray(z, dtype=np.float32) condition_array = self._prepare_condition(condition, latent.shape[0]) dec1_pre = latent @ p["W_dec1"] if c.condition_dim > 0: dec1_pre = dec1_pre + condition_array @ p["W_cond_dec1"] dec1_pre = dec1_pre + p["b_dec1"] dec1 = relu(dec1_pre) dec2_pre = dec1 @ p["W_dec2"] if c.condition_dim > 0: dec2_pre = dec2_pre + condition_array @ p["W_cond_dec2"] dec2_pre = dec2_pre + p["b_dec2"] dec2 = relu(dec2_pre) flat_dec = dec2 @ p["W_dec3"] + p["b_dec3"] tile_dec_pre = flat_dec.reshape(latent.shape[0], c.num_tiles, c.tile_embedding_dim) tile_dec = np.tanh(tile_dec_pre) out_pre = np.einsum("bne,ef->bnf", tile_dec, p["W_tile_out"]) + p["b_tile_out"] if c.condition_dim > 0: out_pre = out_pre + (condition_array @ p["W_cond_out"]).reshape(latent.shape[0], 1, c.tile_feature_dim) return self._decode_output_from_logits(out_pre, training=False)["output"].reshape(latent.shape[0], -1) def sample( self, count: int, temperature: float = 1.0, seed: int | None = None, condition: np.ndarray | None = None, ) -> np.ndarray: rng = np.random.default_rng(self.config.seed if seed is None else seed) if self.config.condition_dim > 0: condition_array = self._prepare_condition(condition, count) prior_mu, prior_logvar = self._conditional_prior(condition_array) prior_std = np.exp(0.5 * np.clip(prior_logvar, -20.0, 20.0)).astype(np.float32) eps = rng.normal(0.0, temperature, size=(count, self.config.latent_dim)).astype(np.float32) z = prior_mu + eps * prior_std else: z = rng.normal(0.0, temperature, size=(count, self.config.latent_dim)).astype(np.float32) return self.decode(z, condition=condition) def save( self, path: str | Path, codec_config: dict[str, Any] | None = None, extra_metadata: dict[str, Any] | None = None, ) -> None: target = Path(path) target.parent.mkdir(parents=True, exist_ok=True) payload: dict[str, Any] = { **self.params, "feature_weights": self.feature_weights.astype(np.float32), "terrain_class_weights": self.terrain_class_weights.astype(np.float32), "config_json": np.array(json.dumps(self.config.to_dict())), "codec_json": np.array(json.dumps(codec_config or {})), "metadata_json": np.array(json.dumps(extra_metadata or {})), } np.savez(target, **payload) @classmethod def load(cls, path: str | Path) -> tuple["NumpyVAE", dict[str, Any], dict[str, Any]]: source = Path(path) data = np.load(source, allow_pickle=False) config = VAEConfig.from_dict(json.loads(str(data["config_json"].item()))) codec_config = json.loads(str(data["codec_json"].item())) metadata = json.loads(str(data["metadata_json"].item())) terrain_class_weights = None if "terrain_class_weights" in data: terrain_class_weights = data["terrain_class_weights"].astype(np.float32) model = cls( config=config, feature_weights=data["feature_weights"].astype(np.float32), terrain_class_weights=terrain_class_weights, ) for name in model.params: if name in data: model.params[name] = data[name].astype(np.float32) elif name.startswith("W_cond") or name.startswith("W_prior") or name.startswith("b_prior"): model.params[name] = np.zeros_like(model.params[name], dtype=np.float32) return model, codec_config, metadata