SyntheticMDProductions's picture
Update ADAM safety, UI, and model workflows (#1)
c61c435
Raw History Blame Contribute Delete
2.48 kB
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)),
)