from __future__ import annotations import torch from torch import nn from torch.nn import functional as F class NeuralCellularAutomata(nn.Module): def __init__(self, channels: int = 16, hidden_size: int = 128, fire_rate: float = 0.5): super().__init__() if channels < 4: raise ValueError("NCA requires at least 4 cell channels (RGBA).") self.channels = channels self.hidden_size = hidden_size self.fire_rate = fire_rate self.update_net = nn.Sequential( nn.Conv2d(channels * 3, hidden_size, kernel_size=1), nn.ReLU(), nn.Conv2d(hidden_size, channels, kernel_size=1, bias=False), ) nn.init.zeros_(self.update_net[-1].weight) identity = torch.tensor([[0.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 0.0]]) sobel_x = torch.tensor([[-1.0, 0.0, 1.0], [-2.0, 0.0, 2.0], [-1.0, 0.0, 1.0]]) / 8.0 sobel_y = sobel_x.t() kernels = torch.stack((identity, sobel_x, sobel_y))[:, None] self.register_buffer("perception_kernels", kernels) def perceive(self, state: torch.Tensor) -> torch.Tensor: kernels = self.perception_kernels.repeat(self.channels, 1, 1, 1) return F.conv2d(state, kernels, padding=1, groups=self.channels) @staticmethod def living_mask(state: torch.Tensor) -> torch.Tensor: return F.max_pool2d(state[:, 3:4], kernel_size=3, stride=1, padding=1) > 0.1 def forward(self, state: torch.Tensor, fire_rate: float | None = None) -> torch.Tensor: pre_life = self.living_mask(state) delta = self.update_net(self.perceive(state)) rate = self.fire_rate if fire_rate is None else fire_rate stochastic = (torch.rand_like(state[:, :1]) <= rate).to(state.dtype) state = state + delta * stochastic post_life = self.living_mask(state) return state * (pre_life & post_life).to(state.dtype) def create_seed(batch_size: int, channels: int, resolution: int, device: torch.device) -> torch.Tensor: state = torch.zeros(batch_size, channels, resolution, resolution, device=device) center = resolution // 2 state[:, 3, center, center] = 1.0 return state def create_model(config: dict) -> NeuralCellularAutomata: return NeuralCellularAutomata( channels=int(config.get("cell_channels", 16)), hidden_size=int(config.get("hidden_size", 128)), fire_rate=float(config.get("fire_rate", 0.5)), )