Download models/neural_cellular_automata/model.py from SyntheticMDProductions/AI_Development_Automation_Manager: direct link, hf CLI and curl.
- Browser
- Download file 2.48 kB
-
https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/main/models/neural_cellular_automata/model.py
- Command line
-
hf download hf://SyntheticMDProductions/AI_Development_Automation_Manager/models/neural_cellular_automata/model.py
-
curl -L -o model.py https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/main/models/neural_cellular_automata/model.py
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) | |
| 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)), | |
| ) | |