ARotting's picture
Publish Neural Gray-Scott reaction-diffusion surrogate
626ab93 verified
Raw
History Blame Contribute Delete
8.36 kB
from __future__ import annotations
import json
import random
from pathlib import Path
import numpy as np
import torch
import trackio
from model import TuringSurrogate, parameter_count
from physics import gray_scott_step, initial_state
from PIL import Image
from safetensors.torch import save_file
from torch.nn import functional as F
from torch.utils.data import DataLoader, TensorDataset
PROJECT_DIR = Path(__file__).resolve().parent
DATA_DIR = PROJECT_DIR / "data"
ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "turing-neural-field"
def seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
def load_data() -> dict[str, torch.Tensor]:
arrays = np.load(DATA_DIR / "gray_scott_trajectories.npz")
return {
key: torch.from_numpy(arrays[key])
for key in ["states", "next_states", "feeds", "kills", "trajectory_ids"]
}
@torch.inference_mode()
def one_step_metrics(
model: TuringSurrogate,
state: torch.Tensor,
target: torch.Tensor,
feed: torch.Tensor,
kill: torch.Tensor,
) -> dict:
model.eval()
predictions = []
for start in range(0, len(state), 64):
predictions.append(
model(
state[start : start + 64],
feed[start : start + 64],
kill[start : start + 64],
)
)
prediction = torch.cat(predictions)
model_mse = F.mse_loss(prediction, target).item()
persistence_mse = F.mse_loss(state, target).item()
return {
"model_mse": model_mse,
"persistence_mse": persistence_mse,
"improvement_percent": 100 * (persistence_mse - model_mse) / persistence_mse,
}
@torch.inference_mode()
def rollout(
model: TuringSurrogate,
steps: int = 120,
) -> tuple[dict, list[torch.Tensor], list[torch.Tensor]]:
feed = torch.tensor([0.042])
kill = torch.tensor([0.061])
physics_state = initial_state(1, 32, seed=4096)
neural_state = physics_state.clone()
physics_frames = [physics_state.clone()]
neural_frames = [neural_state.clone()]
model.eval()
for step in range(1, steps + 1):
physics_state = gray_scott_step(physics_state, feed, kill)
neural_state = model(neural_state, feed, kill)
if step % 5 == 0:
physics_frames.append(physics_state.clone())
neural_frames.append(neural_state.clone())
field_mse = F.mse_loss(neural_state, physics_state).item()
physics_v = physics_state[:, 1]
neural_v = neural_state[:, 1]
centered_physics = physics_v - physics_v.mean()
centered_neural = neural_v - neural_v.mean()
correlation = float(
(centered_physics * centered_neural).sum()
/ (
torch.sqrt((centered_physics.square()).sum())
* torch.sqrt((centered_neural.square()).sum())
+ 1e-9
)
)
metrics = {
"steps": steps,
"final_field_mse": field_mse,
"v_field_correlation": correlation,
"physics_v_mean": float(physics_v.mean()),
"neural_v_mean": float(neural_v.mean()),
"physics_v_spatial_std": float(physics_v.std()),
"neural_v_spatial_std": float(neural_v.std()),
}
return metrics, physics_frames, neural_frames
def field_image(state: torch.Tensor) -> np.ndarray:
v = state[0, 1].cpu().numpy()
normalized = np.clip(v / max(0.05, float(v.max())), 0, 1)
red = np.clip(normalized * 70, 0, 255)
green = np.clip(normalized * 210, 0, 255)
blue = np.clip(30 + normalized * 225, 0, 255)
return np.stack([red, green, blue], axis=2).astype(np.uint8)
def save_comparison_gif(
physics_frames: list[torch.Tensor],
neural_frames: list[torch.Tensor],
path: Path,
) -> None:
images = []
for physics_state, neural_state in zip(
physics_frames,
neural_frames,
strict=True,
):
physics_image = field_image(physics_state)
neural_image = field_image(neural_state)
separator = np.full((32, 2, 3), 245, dtype=np.uint8)
combined = np.concatenate([physics_image, separator, neural_image], axis=1)
images.append(
Image.fromarray(combined).resize((1056, 512), Image.Resampling.NEAREST)
)
images[0].save(
path,
save_all=True,
append_images=images[1:],
duration=120,
loop=0,
)
images[-1].save(path.with_name("final_comparison.png"))
def main() -> None:
seed_everything(2040)
data = load_data()
train_mask = data["trajectory_ids"] < 28
validation_mask = (data["trajectory_ids"] >= 28) & (data["trajectory_ids"] < 32)
test_mask = data["trajectory_ids"] >= 32
train_dataset = TensorDataset(
data["states"][train_mask],
data["next_states"][train_mask],
data["feeds"][train_mask],
data["kills"][train_mask],
)
loader = DataLoader(
train_dataset,
batch_size=32,
shuffle=True,
generator=torch.Generator().manual_seed(2040),
)
model = TuringSurrogate()
optimizer = torch.optim.AdamW(model.parameters(), lr=0.0015, weight_decay=0.001)
epochs = 55
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)
best_validation_mse = float("inf")
best_epoch = 0
best_state = None
trackio.init(
project="turing-neural-field",
name="gray-scott-surrogate-v1",
config={
"parameters": parameter_count(model),
"train_trajectories": 28,
"validation_trajectories": 4,
"test_trajectories": 4,
"epochs": epochs,
},
)
for epoch in range(1, epochs + 1):
model.train()
losses = []
for state, target, feed, kill in loader:
prediction = model(state, feed, kill)
loss = F.mse_loss(prediction, target) * 10_000
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
losses.append(loss.item())
scheduler.step()
validation = one_step_metrics(
model,
data["states"][validation_mask],
data["next_states"][validation_mask],
data["feeds"][validation_mask],
data["kills"][validation_mask],
)
if validation["model_mse"] < best_validation_mse:
best_validation_mse = validation["model_mse"]
best_epoch = epoch
best_state = {
key: value.detach().cpu().clone()
for key, value in model.state_dict().items()
}
trackio.log(
{
"epoch": epoch,
"train_scaled_mse": float(np.mean(losses)),
"validation_one_step_mse": validation["model_mse"],
"learning_rate": scheduler.get_last_lr()[0],
}
)
assert best_state is not None
model.load_state_dict(best_state)
test = one_step_metrics(
model,
data["states"][test_mask],
data["next_states"][test_mask],
data["feeds"][test_mask],
data["kills"][test_mask],
)
rollout_metrics, physics_frames, neural_frames = rollout(model)
results = {
"model": "Turing Neural Field",
"parameters": parameter_count(model),
"best_epoch": best_epoch,
"test_states": int(test_mask.sum()),
"one_step_test": test,
"held_out_rollout": rollout_metrics,
"gif_layout": "left physics simulator, right neural surrogate",
}
trackio.log(
{
"test_one_step_mse": test["model_mse"],
"rollout_final_field_mse": rollout_metrics["final_field_mse"],
"rollout_v_correlation": rollout_metrics["v_field_correlation"],
}
)
trackio.finish()
ARTIFACT_DIR.mkdir(parents=True, exist_ok=True)
save_file(model.state_dict(), ARTIFACT_DIR / "model.safetensors")
save_comparison_gif(
physics_frames,
neural_frames,
ARTIFACT_DIR / "physics_vs_neural.gif",
)
(ARTIFACT_DIR / "evaluation.json").write_text(
json.dumps(results, indent=2),
encoding="utf-8",
)
print(json.dumps(results, indent=2))
if __name__ == "__main__":
main()