mofscreen-agent-api / coprinet /pricePrediction /nets /basicArchitectures.py
andy88836's picture
Deploy MOFScreen-Agent FastAPI backend
3de59dd verified
Raw
History Blame Contribute Delete
10.6 kB
import torch
import numpy as np
from torch import nn
import torch.nn.functional as F
from torch_geometric.nn import BatchNorm, global_add_pool, AttentiveFP, Set2Set
from torch_geometric.nn.conv import MessagePassing
from torch_geometric.nn.dense.linear import Linear as GeometricLinear
from torch_geometric.utils import degree, scatter
from torch.nn import ModuleList, Linear, ReLU, Sequential, Dropout
class LegacyPNAConv(MessagePassing):
"""PyG 2.0.1 PNA behavior used to train the bundled checkpoint."""
def __init__(self, in_channels, out_channels, aggregators, scalers, deg,
edge_dim=None, towers=1, pre_layers=1, post_layers=1,
divide_input=False, **kwargs):
kwargs.setdefault('aggr', None)
super().__init__(node_dim=0, **kwargs)
if divide_input:
assert in_channels % towers == 0
assert out_channels % towers == 0
self.in_channels = in_channels
self.out_channels = out_channels
self.aggregators = aggregators
self.scalers = scalers
self.edge_dim = edge_dim
self.towers = towers
self.divide_input = divide_input
self.F_in = in_channels // towers if divide_input else in_channels
self.F_out = out_channels // towers
deg = deg.to(torch.float)
self.avg_deg = {
'lin': deg.mean().item(),
'log': (deg + 1).log().mean().item(),
'exp': deg.exp().mean().item(),
}
if edge_dim is not None:
self.edge_encoder = GeometricLinear(edge_dim, self.F_in)
self.pre_nns = ModuleList()
self.post_nns = ModuleList()
for _ in range(towers):
modules = [GeometricLinear((3 if edge_dim else 2) * self.F_in, self.F_in)]
for _ in range(pre_layers - 1):
modules += [ReLU(), GeometricLinear(self.F_in, self.F_in)]
self.pre_nns.append(Sequential(*modules))
post_in = (len(aggregators) * len(scalers) + 1) * self.F_in
modules = [GeometricLinear(post_in, self.F_out)]
for _ in range(post_layers - 1):
modules += [ReLU(), GeometricLinear(self.F_out, self.F_out)]
self.post_nns.append(Sequential(*modules))
self.lin = GeometricLinear(out_channels, out_channels)
self.reset_parameters()
def reset_parameters(self):
if self.edge_dim is not None:
self.edge_encoder.reset_parameters()
for network in self.pre_nns:
for module in network:
if hasattr(module, 'reset_parameters'):
module.reset_parameters()
for network in self.post_nns:
for module in network:
if hasattr(module, 'reset_parameters'):
module.reset_parameters()
self.lin.reset_parameters()
def forward(self, x, edge_index, edge_attr=None):
if self.divide_input:
x = x.view(-1, self.towers, self.F_in)
else:
x = x.view(-1, 1, self.F_in).repeat(1, self.towers, 1)
out = self.propagate(edge_index, x=x, edge_attr=edge_attr, size=None)
out = torch.cat([x, out], dim=-1)
out = torch.cat([network(out[:, i]) for i, network in enumerate(self.post_nns)], dim=1)
return self.lin(out)
def message(self, x_i, x_j, edge_attr):
if edge_attr is not None:
edge_attr = self.edge_encoder(edge_attr)
edge_attr = edge_attr.view(-1, 1, self.F_in).repeat(1, self.towers, 1)
features = torch.cat([x_i, x_j, edge_attr], dim=-1)
else:
features = torch.cat([x_i, x_j], dim=-1)
return torch.stack([network(features[:, i]) for i, network in enumerate(self.pre_nns)], dim=1)
def aggregate(self, inputs, index, dim_size=None):
outputs = []
for aggregator in self.aggregators:
if aggregator in {'sum', 'mean', 'min', 'max'}:
output = scatter(inputs, index, dim=0, dim_size=dim_size, reduce=aggregator)
elif aggregator in {'var', 'std'}:
mean = scatter(inputs, index, dim=0, dim_size=dim_size, reduce='mean')
mean_squares = scatter(inputs * inputs, index, dim=0, dim_size=dim_size, reduce='mean')
output = mean_squares - mean * mean
if aggregator == 'std':
output = torch.sqrt(torch.relu(output) + 1e-5)
else:
raise ValueError(f'Unknown aggregator "{aggregator}"')
outputs.append(output)
output = torch.cat(outputs, dim=-1)
node_degree = degree(index, dim_size, dtype=inputs.dtype).clamp_(1).view(-1, 1, 1)
outputs = []
for scaler in self.scalers:
if scaler == 'identity':
pass
elif scaler == 'amplification':
output = output * (torch.log(node_degree + 1) / self.avg_deg['log'])
elif scaler == 'attenuation':
output = output * (self.avg_deg['log'] / torch.log(node_degree + 1))
elif scaler == 'linear':
output = output * (node_degree / self.avg_deg['lin'])
elif scaler == 'inverse_linear':
output = output * (self.avg_deg['lin'] / node_degree)
else:
raise ValueError(f'Unknown scaler "{scaler}"')
outputs.append(output)
return torch.cat(outputs, dim=-1)
class MLP(nn.Module):
def __init__(self,dims, n_layers, hidden_size, dropout=0 ):
super().__init__()
self.n_layers = n_layers
self.hidden_size = hidden_size
self.dims = dims
self.dropout = dropout
def block(in_size, n_hidden):
layers = [
nn.Linear(in_size, n_hidden),
nn.BatchNorm1d(n_hidden),
nn.ReLU(),
]
if self.dropout > 0:
layers.append(
nn.Dropout(self.dropout),
)
return layers
# Define PyTorch model
self.model = nn.Sequential(
*block(np.prod(dims), self.hidden_size)
)
self.latent_size = self.hidden_size
def forward(self, x):
return self.model(x)
class GNN_PNAConv(torch.nn.Module):
def __init__(self,
nodes_n_features: int,
edges_n_features: int,
deg: torch.tensor ,
n_layers: int = 6,
hidden_size_node: int = 75,
hidden_size_edges: int = 50,
towers: int = 5,
fcc_hidden_size: int = 50,
dropout: float = 0,
use_fds: bool = False,
**args
):
super(GNN_PNAConv, self).__init__()
self.node_emb = MLP(nodes_n_features, 1, hidden_size_node)
self.edge_emb = MLP(edges_n_features, 1, hidden_size_edges)
aggregators = ['mean', 'min', 'max', 'std']
scalers = ['identity', 'amplification', 'attenuation']
self.convs = ModuleList()
self.batch_norms = ModuleList()
for _ in range(n_layers):
conv = LegacyPNAConv(in_channels=hidden_size_node, out_channels=hidden_size_node,
aggregators=aggregators, scalers=scalers, deg=deg,
edge_dim=hidden_size_edges, towers=towers, pre_layers=2, post_layers=2,
divide_input=False)
self.convs.append(conv)
self.batch_norms.append(BatchNorm(hidden_size_node))
self.set2set = Set2Set(hidden_size_node, processing_steps=6)
fc_layers = [
Linear(2*hidden_size_node, hidden_size_node), BatchNorm(hidden_size_node), ReLU(),
Linear(hidden_size_node, fcc_hidden_size), BatchNorm(fcc_hidden_size), ReLU(),
]
self.fcc_hidden_size = fcc_hidden_size
if dropout>0:
fc_layers += [ Dropout(p=dropout) ]
self.mlp = Sequential( *fc_layers )
def forward(self, x, edge_index, edge_attr, batch):
x = self.node_emb(x.squeeze())
edge_attr = self.edge_emb(edge_attr)
for conv, batch_norm in zip(self.convs, self.batch_norms):
x = F.relu(batch_norm(conv(x, edge_index, edge_attr)))
x = self.set2set(x, batch) #Set2Set #GlobalAttention #GraphMultisetTransformer
return self.mlp(x)
class QdolarAR(torch.nn.Module):
def __init__(self,
nodes_n_features: int,
*args, **kargs
):
super(QdolarAR, self).__init__()
n_layers: int = 5
hidden_size_node: int = nodes_n_features
fcc_hidden_size: int = 100
dropout: float = 0.28
self.fcc_hidden_size = fcc_hidden_size
self.dropout = dropout
self.n_layers = n_layers
self.hidden_size_node = hidden_size_node
assert n_layers >=3
fc_layers = [
BatchNorm(hidden_size_node),
Linear( hidden_size_node, fcc_hidden_size), BatchNorm(fcc_hidden_size), ReLU(),
]
for i in range(n_layers-2):
fc_layers += [
Linear(fcc_hidden_size, fcc_hidden_size), BatchNorm(fcc_hidden_size), ReLU(),
]
if dropout>0:
fc_layers += [ Dropout(p=dropout) ]
self.mlp = Sequential( *fc_layers )
def forward(self, x, edge_index, edge_attr, batch):
assert len(batch) == len(batch.unique())
assert not torch.isnan(x).any(), "Error, x contains nan"
return self.mlp(x)
class GNN_AttentiveFP(torch.nn.Module):
def __init__(self,
nodes_n_features: int,
edges_n_features: int,
deg: torch.tensor ,
n_layers: int = 4,
hidden_size_node: int = 75,
hidden_size_edges: int = 50,
towers: int = 5,
fcc_hidden_size: int = 50,
dropout: float = 0,
use_fds: bool = False,
**args
):
super(GNN_AttentiveFP, self).__init__()
self.attent_net = AttentiveFP(nodes_n_features, hidden_size_node, hidden_size_node, edges_n_features,
n_layers, n_layers, dropout= dropout)
fc_layers = [
Linear(hidden_size_node, hidden_size_edges), ReLU(),
Linear(hidden_size_edges, fcc_hidden_size), ReLU(),
]
self.fcc_hidden_size = fcc_hidden_size
if dropout>0:
fc_layers += Dropout(p=dropout)
self.mlp = Sequential( *fc_layers )
def forward(self, x, edge_index, edge_attr, batch):
x = self.attent_net(x, edge_index, edge_attr, batch)
return self.mlp(x)