import torch import torch.nn as nn from e3nn import o3 from e3nn.nn import GatedMLP
class EquivariantNetwork(nn.Module): """等变神经网络""" def __init__(self, irreps_in, irreps_hidden, irreps_out, num_layers=4): super().__init__() self.layers = nn.ModuleList() irreps_current = irreps_in for _ in range(num_layers): layer = o3.Convolution( irreps_current, irreps_hidden, irreps_hidden ) self.layers.append(layer) irreps_current = irreps_hidden self.output = o3.Linear(irreps_hidden, irreps_out) def forward(self, x, edge_index, edge_attr, t): """等变前向传播""" t_embed = self.time_embed(t) for layer in self.layers: x = layer(x, edge_index, edge_attr) x = x + t_embed return self.output(x)
class EDM(nn.Module): """等变扩散模型""" def __init__(self, num_atom_types=100, hidden_dim=64, num_steps=1000): super().__init__() self.num_atom_types = num_atom_types self.num_steps = num_steps self.coordinate_net = EquivariantNetwork( irreps_in='1x1o', irreps_hidden=f'{hidden_dim}x0e+{hidden_dim//2}x1o', irreps_out='1x1o' ) self.atom_type_net = nn.Sequential( nn.Linear(hidden_dim, hidden_dim), nn.SiLU(), nn.Linear(hidden_dim, num_atom_types) ) self.time_embed = nn.Sequential( nn.Linear(1, hidden_dim), nn.SiLU(), nn.Linear(hidden_dim, hidden_dim) ) self.betas = torch.linspace(1e-4, 0.02, num_steps) self.alphas = 1 - self.betas self.alpha_bar = torch.cumprod(self.alphas, dim=0) def forward_diffusion(self, x0, t): """前向扩散过程""" noise = torch.randn_like(x0) alpha_bar_t = self.alpha_bar[t].view(-1, 1, 1) xt = torch.sqrt(alpha_bar_t) * x0 + torch.sqrt(1 - alpha_bar_t) * noise return xt, noise def reverse_step(self, xt, t, edge_index): """反向去噪步骤""" t_embed = self.time_embed(t.view(-1, 1).float()) noise_pred = self.coordinate_net(xt, edge_index, None, t_embed) alpha_t = self.alphas[t].view(-1, 1, 1) alpha_bar_t = self.alpha_bar[t].view(-1, 1, 1) mean = (xt - (1 - alpha_t) / torch.sqrt(1 - alpha_bar_t) * noise_pred) mean = mean / torch.sqrt(alpha_t) if self.training: noise = torch.randn_like(xt) std = torch.sqrt((1 - alpha_bar_t) / (1 - alpha_bar_t + 1e-8)) return mean + std * noise else: return mean def generate(self, num_molecules, num_atoms, edge_index): """生成分子""" x = torch.randn(num_molecules, num_atoms, 3) for t in reversed(range(self.num_steps)): t_tensor = torch.full((num_molecules,), t, device=x.device) x = self.reverse_step(x, t_tensor, edge_index) return x def training_loss(self, x0, edge_index): """计算训练损失""" batch_size = x0.shape[0] t = torch.randint(0, self.num_steps, (batch_size,), device=x0.device) xt, noise = self.forward_diffusion(x0, t) t_embed = self.time_embed(t.view(-1, 1).float()) noise_pred = self.coordinate_net(xt, edge_index, None, t_embed) loss = nn.MSELoss()(noise_pred, noise) return loss
def train_edm(model, train_loader, epochs=100): """训练EDM模型""" optimizer = torch.optim.Adam(model.parameters(), lr=1e-4) for epoch in range(epochs): total_loss = 0 for batch in train_loader: loss = model.training_loss( batch['positions'], batch['edge_index'] ) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() total_loss += loss.item() if (epoch + 1) % 10 == 0: print(f"Epoch {epoch+1}, Loss: {total_loss:.6f}")
|