sankalpsthakur's picture
Publish clean-room generator, training, and export source
e6e6c97 verified
Raw
History Blame Contribute Delete
1.08 kB
"""Tiny multi-output PyTorch surrogate with embedded normalization."""
from __future__ import annotations
import torch
class PumpSurrogate(torch.nn.Module):
def __init__(
self,
input_mean: torch.Tensor,
input_std: torch.Tensor,
output_mean: torch.Tensor,
output_std: torch.Tensor,
) -> None:
super().__init__()
self.register_buffer("input_mean", input_mean.float())
self.register_buffer("input_std", input_std.float())
self.register_buffer("output_mean", output_mean.float())
self.register_buffer("output_std", output_std.float())
self.core = torch.nn.Sequential(
torch.nn.Linear(6, 32),
torch.nn.ReLU(),
torch.nn.Linear(32, 32),
torch.nn.ReLU(),
torch.nn.Linear(32, 6),
)
def forward(self, features: torch.Tensor) -> torch.Tensor:
normalized = (features - self.input_mean) / self.input_std
prediction = self.core(normalized)
return prediction * self.output_std + self.output_mean