BlidReview's picture
weights, code, eval script
bdce880 verified
Raw
History Blame Contribute Delete
6.1 kB
import torch
import torch.nn as nn
from timm.models.layers import trunc_normal_
from model.Embedding import timestep_embedding
import numpy as np
from model.Physics_Attention import Physics_Attention_Irregular_Mesh
ACTIVATION = {'gelu': nn.GELU, 'tanh': nn.Tanh, 'sigmoid': nn.Sigmoid, 'relu': nn.ReLU, 'leaky_relu': nn.LeakyReLU(0.1),
'softplus': nn.Softplus, 'ELU': nn.ELU, 'silu': nn.SiLU}
class MLP(nn.Module):
def __init__(self, n_input, n_hidden, n_output, n_layers=1, act='gelu', res=True):
super(MLP, self).__init__()
if act in ACTIVATION.keys():
act = ACTIVATION[act]
else:
raise NotImplementedError
self.n_input = n_input
self.n_hidden = n_hidden
self.n_output = n_output
self.n_layers = n_layers
self.res = res
self.linear_pre = nn.Sequential(nn.Linear(n_input, n_hidden), act())
self.linear_post = nn.Linear(n_hidden, n_output)
self.linears = nn.ModuleList([nn.Sequential(nn.Linear(n_hidden, n_hidden), act()) for _ in range(n_layers)])
def forward(self, x):
x = self.linear_pre(x)
for i in range(self.n_layers):
if self.res:
x = self.linears[i](x) + x
else:
x = self.linears[i](x)
x = self.linear_post(x)
return x
class Transolver_block(nn.Module):
"""Transformer encoder block."""
def __init__(
self,
num_heads: int,
hidden_dim: int,
dropout: float,
act='gelu',
mlp_ratio=4,
last_layer=False,
out_dim=1,
slice_num=32,
):
super().__init__()
self.last_layer = last_layer
self.ln_1 = nn.LayerNorm(hidden_dim)
self.Attn = Physics_Attention_Irregular_Mesh(hidden_dim, heads=num_heads, dim_head=hidden_dim // num_heads,
dropout=dropout, slice_num=slice_num)
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 self.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))
else:
return fx
class Model(nn.Module):
def __init__(self,
space_dim=1,
n_layers=5,
n_hidden=256,
dropout=0.0,
n_head=8,
Time_Input=False,
act='gelu',
mlp_ratio=1,
fun_dim=1,
out_dim=1,
slice_num=32,
ref=8,
unified_pos=False
):
super(Model, self).__init__()
self.__name__ = 'Transolver_1D'
self.ref = ref
self.unified_pos = unified_pos
self.Time_Input = Time_Input
self.n_hidden = n_hidden
self.space_dim = space_dim
if self.unified_pos:
self.preprocess = MLP(fun_dim + self.ref * self.ref, n_hidden * 2, n_hidden, n_layers=0, res=False, act=act)
else:
self.preprocess = MLP(fun_dim + space_dim, n_hidden * 2, n_hidden, n_layers=0, res=False, act=act)
if Time_Input:
self.time_fc = nn.Sequential(nn.Linear(n_hidden, n_hidden), nn.SiLU(), nn.Linear(n_hidden, n_hidden))
self.blocks = nn.ModuleList([Transolver_block(num_heads=n_head, hidden_dim=n_hidden,
dropout=dropout,
act=act,
mlp_ratio=mlp_ratio,
out_dim=out_dim,
slice_num=slice_num,
last_layer=(_ == n_layers - 1))
for _ in range(n_layers)])
self.initialize_weights()
self.placeholder = nn.Parameter((1 / (n_hidden)) * torch.rand(n_hidden, dtype=torch.float))
def initialize_weights(self):
self.apply(self._init_weights)
def _init_weights(self, m):
if isinstance(m, nn.Linear):
trunc_normal_(m.weight, std=0.02)
if isinstance(m, nn.Linear) and m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, (nn.LayerNorm, nn.BatchNorm1d)):
nn.init.constant_(m.bias, 0)
nn.init.constant_(m.weight, 1.0)
def get_grid(self, x, batchsize=1):
# x: B N 2
# grid_ref
gridx = torch.tensor(np.linspace(0, 1, self.ref), dtype=torch.float)
gridx = gridx.reshape(1, self.ref, 1, 1).repeat([batchsize, 1, self.ref, 1])
gridy = torch.tensor(np.linspace(0, 1, self.ref), dtype=torch.float)
gridy = gridy.reshape(1, 1, self.ref, 1).repeat([batchsize, self.ref, 1, 1])
grid_ref = torch.cat((gridx, gridy), dim=-1).cuda().reshape(batchsize, self.ref * self.ref, 2) # B H W 8 8 2
pos = torch.sqrt(torch.sum((x[:, :, None, :] - grid_ref[:, None, :, :]) ** 2, dim=-1)). \
reshape(batchsize, x.shape[1], self.ref * self.ref).contiguous()
return pos
def forward(self, x, fx, T=None):
if self.unified_pos:
x = self.get_grid(x, x.shape[0])
if fx is not None:
fx = torch.cat((x, fx), -1)
fx = self.preprocess(fx)
else:
fx = self.preprocess(x)
fx = fx + self.placeholder[None, None, :]
if T is not None:
Time_emb = timestep_embedding(T, self.n_hidden).repeat(1, x.shape[1], 1)
Time_emb = self.time_fc(Time_emb)
fx = fx + Time_emb
for block in self.blocks:
fx = block(fx)
return fx