AerovigilAI / model.py
rajaram-90's picture
Create model.py
b195a41 verified
Raw
History Blame Contribute Delete
5.05 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class BayesianLinear(nn.Module):
"""
Variational Bayesian Linear Layer using the Reparameterization Trick.
"""
def __init__(self, in_features: int, out_features: int, prior_sigma: float = 1.0):
super(BayesianLinear, self).__init__()
self.in_features = in_features
self.out_features = out_features
self.prior_sigma = prior_sigma
# Weight variational parameters (Mean and Rho)
self.weight_mu = nn.Parameter(torch.Tensor(out_features, in_features).normal_(0, 0.1))
self.weight_rho = nn.Parameter(torch.Tensor(out_features, in_features).fill_(-3.0))
# Bias variational parameters
self.bias_mu = nn.Parameter(torch.Tensor(out_features).normal_(0, 0.1))
self.bias_rho = nn.Parameter(torch.Tensor(out_features).fill_(-3.0))
def forward(self, x: torch.Tensor) -> torch.Tensor:
# Reparameterization trick: std = log(1 + exp(rho))
weight_sigma = torch.log1p(torch.exp(self.weight_rho))
bias_sigma = torch.log1p(torch.exp(self.bias_rho))
epsilon_w = torch.randn_like(self.weight_mu)
epsilon_b = torch.randn_like(self.bias_mu)
weight = self.weight_mu + weight_sigma * epsilon_w
bias = self.bias_mu + bias_sigma * epsilon_b
return F.linear(x, weight, bias)
def kl_divergence(self) -> torch.Tensor:
"""
Calculates KL Divergence between variational posterior q(w) and Gaussian prior p(w).
"""
weight_sigma = torch.log1p(torch.exp(self.weight_rho))
bias_sigma = torch.log1p(torch.exp(self.bias_rho))
kl_w = torch.sum(
torch.log(self.prior_sigma / weight_sigma) +
(weight_sigma**2 + self.weight_mu**2) / (2 * self.prior_sigma**2) - 0.5
)
kl_b = torch.sum(
torch.log(self.prior_sigma / bias_sigma) +
(bias_sigma**2 + self.bias_mu**2) / (2 * self.prior_sigma**2) - 0.5
)
return kl_w + kl_b
class PhysicsGuidedBNN(nn.Module):
"""
Physics-Guided Bayesian Neural Network for Wind Turbine Diagnostics.
"""
def __init__(self, config: dict):
super(PhysicsGuidedBNN, self).__init__()
self.config = config
in_dim = config["in_features"]
hidden_dims = config["hidden_dims"]
num_classes = config["num_classes"]
layers = []
curr_dim = in_dim
for h_dim in hidden_dims:
layers.append(BayesianLinear(curr_dim, h_dim, prior_sigma=config.get("prior_sigma", 1.0)))
layers.append(nn.ReLU())
layers.append(nn.BatchNorm1d(h_dim))
curr_dim = h_dim
self.backbone = nn.Sequential(*layers)
self.classifier = BayesianLinear(curr_dim, num_classes, prior_sigma=config.get("prior_sigma", 1.0))
def forward(self, x: torch.Tensor) -> torch.Tensor:
features = self.backbone(x)
logits = self.classifier(features)
return logits
def total_kl_divergence(self) -> torch.Tensor:
kl = torch.tensor(0.0, device=next(self.parameters()).device)
for module in self.modules():
if isinstance(module, BayesianLinear):
kl = kl + module.kl_divergence()
return kl
def compute_physics_loss(self, x: torch.Tensor, logits: torch.Tensor) -> torch.Tensor:
"""
Computes physical consistency violation penalty.
Assumptions on feature indices:
x[:, 0] -> Wind Speed (m/s)
x[:, 1] -> Generated Power (kW)
x[:, 2] -> Gearbox Oil Temperature (°C)
x[:, 3] -> Generator Speed (RPM)
"""
probs = F.softmax(logits, dim=-1)[:, 1] # Failure probability
wind_speed = x[:, 0]
power = x[:, 1]
gearbox_temp = x[:, 2]
gen_speed = x[:, 3]
# 1. Aerodynamic Power Limit (Betz Law Constraint): P <= 0.5 * rho * A * v^3 * Cp_max
rho = self.config["physics_params"]["air_density"]
radius = self.config["physics_params"]["rotor_radius"]
cp_max = self.config["physics_params"]["cp_max"]
area = math.pi * (radius ** 2)
max_theoretical_power = (0.5 * rho * area * (torch.clamp(wind_speed, min=0) ** 3) * cp_max) / 1000.0 # kW
# Power exceeding physical limit violates physics
power_violation = F.relu(power - max_theoretical_power)
# 2. Thermal Energy Dissipation Residual: dT/dt ~ alpha * Power - beta * (T_gearbox - T_ambient)
alpha = self.config["physics_params"]["thermal_alpha"]
beta = self.config["physics_params"]["thermal_beta"]
expected_thermal_load = alpha * power - beta * gearbox_temp
# Unrealistic low temp under high power/speed triggers high physical residual
thermal_residual = F.relu(expected_thermal_load - gearbox_temp)
# Total Physics Regularization Penalty
physics_loss = torch.mean(power_violation + thermal_residual)
return physics_loss