"""Building blocks: MLP and the pre-norm Transolver block. Faithful to Transolver ``model/Transolver_Irregular_Mesh.py`` (MIT). The only structural change is that the attention module is injected (so Stage 2 can swap Physics-Attention for LinearNO without touching anything else). """ import torch.nn as nn ACTIVATION = { "gelu": nn.GELU, "tanh": nn.Tanh, "sigmoid": nn.Sigmoid, "relu": nn.ReLU, "leaky_relu": lambda: nn.LeakyReLU(0.1), "softplus": nn.Softplus, "ELU": nn.ELU, "silu": nn.SiLU, } class MLP(nn.Module): """Transolver MLP: Linear->act (->[Linear->act]^n_layers) ->Linear, optional residual.""" def __init__(self, n_input, n_hidden, n_output, n_layers=1, act="gelu", res=True): super().__init__() if act not in ACTIVATION: raise NotImplementedError(act) act_cls = ACTIVATION[act] self.n_layers = n_layers self.res = res self.linear_pre = nn.Sequential(nn.Linear(n_input, n_hidden), act_cls()) self.linear_post = nn.Linear(n_hidden, n_output) self.linears = nn.ModuleList( [nn.Sequential(nn.Linear(n_hidden, n_hidden), act_cls()) for _ in range(n_layers)] ) def forward(self, x): x = self.linear_pre(x) for layer in self.linears: x = layer(x) + x if self.res else layer(x) return self.linear_post(x) class TransolverBlock(nn.Module): """Pre-norm transformer block; the last block also carries the decoder head. fx = fx + Attn(LayerNorm(fx)) fx = fx + MLP(LayerNorm(fx)) (last block) return Linear(LayerNorm(fx)) -> out_dim """ def __init__( self, attention: nn.Module, hidden_dim: int, dropout: float = 0.0, act: str = "gelu", mlp_ratio: int = 1, last_layer: bool = False, out_dim: int = 1, ): super().__init__() self.last_layer = last_layer self.ln_1 = nn.LayerNorm(hidden_dim) self.Attn = attention self.ln_2 = nn.LayerNorm(hidden_dim) self.mlp = MLP(hidden_dim, hidden_dim * mlp_ratio, hidden_dim, n_layers=0, res=False, act=act) if last_layer: self.ln_3 = nn.LayerNorm(hidden_dim) self.mlp2 = nn.Linear(hidden_dim, out_dim) def forward(self, fx): fx = self.Attn(self.ln_1(fx)) + fx fx = self.mlp(self.ln_2(fx)) + fx if self.last_layer: return self.mlp2(self.ln_3(fx)) return fx