1313from torch_geometric .utils import dense_to_sparse , to_dense_adj
1414from torch_scatter import scatter_add , scatter_mean
1515from torch_sparse import SparseTensor , coalesce
16- from tqdm .auto import tqdm
1716
1817from ..configuration_utils import ConfigMixin
1918from ..modeling_utils import ModelMixin
@@ -460,35 +459,6 @@ def eq_transform(score_d, pos, edge_index, edge_length):
460459 return score_pos
461460
462461
463- def get_beta_schedule (beta_schedule , * , beta_start , beta_end , num_diffusion_timesteps ):
464- def sigmoid (x ):
465- return 1 / (np .exp (- x ) + 1 )
466-
467- if beta_schedule == "quad" :
468- betas = (
469- np .linspace (
470- beta_start ** 0.5 ,
471- beta_end ** 0.5 ,
472- num_diffusion_timesteps ,
473- dtype = np .float64 ,
474- )
475- ** 2
476- )
477- elif beta_schedule == "linear" :
478- betas = np .linspace (beta_start , beta_end , num_diffusion_timesteps , dtype = np .float64 )
479- elif beta_schedule == "const" :
480- betas = beta_end * np .ones (num_diffusion_timesteps , dtype = np .float64 )
481- elif beta_schedule == "jsd" : # 1/T, 1/(T-1), 1/(T-2), ..., 1
482- betas = 1.0 / np .linspace (num_diffusion_timesteps , 1 , num_diffusion_timesteps , dtype = np .float64 )
483- elif beta_schedule == "sigmoid" :
484- betas = np .linspace (- 6 , 6 , num_diffusion_timesteps )
485- betas = sigmoid (betas ) * (beta_end - beta_start ) + beta_start
486- else :
487- raise NotImplementedError (beta_schedule )
488- assert betas .shape == (num_diffusion_timesteps ,)
489- return betas
490-
491-
492462class DualEncoderEpsNetwork (ModelMixin , ConfigMixin ):
493463 def __init__ (
494464 self ,
@@ -497,10 +467,6 @@ def __init__(
497467 num_convs_local ,
498468 cutoff ,
499469 mlp_act ,
500- beta_schedule ,
501- beta_start ,
502- beta_end ,
503- num_diffusion_timesteps ,
504470 edge_order ,
505471 edge_encoder ,
506472 smooth_conv ,
@@ -551,31 +517,14 @@ def __init__(
551517 self .model_global = nn .ModuleList ([self .edge_encoder_global , self .encoder_global , self .grad_global_dist_mlp ])
552518 self .model_local = nn .ModuleList ([self .edge_encoder_local , self .encoder_local , self .grad_local_dist_mlp ])
553519
554- self .model_type = type # type # 'diffusion'; 'dsm'
555-
556- # denoising diffusion
557- ## betas
558- betas = get_beta_schedule (
559- beta_schedule = beta_schedule ,
560- beta_start = beta_start ,
561- beta_end = beta_end ,
562- num_diffusion_timesteps = num_diffusion_timesteps ,
563- )
564- betas = torch .from_numpy (betas ).float ()
565- self .betas = nn .Parameter (betas , requires_grad = False )
566- ## variances
567- alphas = (1.0 - betas ).cumprod (dim = 0 )
568- self .alphas = nn .Parameter (alphas , requires_grad = False )
569- self .num_timesteps = self .betas .size (0 )
570-
571520 def forward (
572521 self ,
573522 atom_type ,
574523 pos ,
575524 bond_index ,
576525 bond_type ,
577526 batch ,
578- time_step ,
527+ time_step , # NOTE, model trained without timestep performed best
579528 edge_index = None ,
580529 edge_type = None ,
581530 edge_length = None ,
@@ -614,7 +563,6 @@ def forward(
614563
615564 # Encoding global
616565 edge_attr_global = self .edge_encoder_global (edge_length = edge_length , edge_type = edge_type ) # Embed edges
617- # edge_attr += temb_edge
618566
619567 # Global
620568 node_attr_global = self .encoder_global (
@@ -648,6 +596,7 @@ def forward(
648596 edge_index = edge_index [:, local_edge_mask ],
649597 edge_attr = edge_attr_local [local_edge_mask ],
650598 ) # (E_local, 2H)
599+
651600 ## Invariant features of edges (bond graph, local)
652601 if isinstance (sigma_edge , torch .Tensor ):
653602 edge_inv_local = self .grad_local_dist_mlp (h_pair_local ) * (
@@ -661,97 +610,6 @@ def forward(
661610 else :
662611 return edge_inv_global , edge_inv_local
663612
664- def get_loss (
665- self ,
666- atom_type ,
667- pos ,
668- bond_index ,
669- bond_type ,
670- batch ,
671- num_nodes_per_graph ,
672- num_graphs ,
673- anneal_power = 2.0 ,
674- return_unreduced_loss = False ,
675- return_unreduced_edge_loss = False ,
676- extend_order = True ,
677- extend_radius = True ,
678- is_sidechain = None ,
679- ):
680- N = atom_type .size (0 )
681- node2graph = batch
682-
683- # Four elements for DDPM: original_data(pos), gaussian_noise(pos_noise), beta(sigma), time_step
684- # Sample noise levels
685- time_step = torch .randint (0 , self .num_timesteps , size = (num_graphs // 2 + 1 ,), device = pos .device )
686- time_step = torch .cat ([time_step , self .num_timesteps - time_step - 1 ], dim = 0 )[:num_graphs ]
687- a = self .alphas .index_select (0 , time_step ) # (G, )
688- # Perterb pos
689- a_pos = a .index_select (0 , node2graph ).unsqueeze (- 1 ) # (N, 1)
690- pos_noise = torch .zeros (size = pos .size (), device = pos .device )
691- pos_noise .normal_ ()
692- pos_perturbed = pos + pos_noise * (1.0 - a_pos ).sqrt () / a_pos .sqrt ()
693-
694- # Update invariant edge features, as shown in equation 5-7
695- edge_inv_global , edge_inv_local , edge_index , edge_type , edge_length , local_edge_mask = self (
696- atom_type = atom_type ,
697- pos = pos_perturbed ,
698- bond_index = bond_index ,
699- bond_type = bond_type ,
700- batch = batch ,
701- time_step = time_step ,
702- return_edges = True ,
703- extend_order = extend_order ,
704- extend_radius = extend_radius ,
705- is_sidechain = is_sidechain ,
706- ) # (E_global, 1), (E_local, 1)
707-
708- edge2graph = node2graph .index_select (0 , edge_index [0 ])
709- # Compute sigmas_edge
710- a_edge = a .index_select (0 , edge2graph ).unsqueeze (- 1 ) # (E, 1)
711-
712- # Compute original and perturbed distances
713- d_gt = get_distance (pos , edge_index ).unsqueeze (- 1 ) # (E, 1)
714- d_perturbed = edge_length
715- # Filtering for protein
716- train_edge_mask = is_train_edge (edge_index , is_sidechain )
717- d_perturbed = torch .where (train_edge_mask .unsqueeze (- 1 ), d_perturbed , d_gt )
718-
719- if self .edge_encoder == "gaussian" :
720- # Distances must be greater than 0
721- d_sgn = torch .sign (d_perturbed )
722- d_perturbed = torch .clamp (d_perturbed * d_sgn , min = 0.01 , max = float ("inf" ))
723- d_target = (d_gt - d_perturbed ) / (1.0 - a_edge ).sqrt () * a_edge .sqrt () # (E_global, 1), denoising direction
724-
725- global_mask = torch .logical_and (
726- torch .logical_or (d_perturbed <= self .cutoff , local_edge_mask .unsqueeze (- 1 )),
727- ~ local_edge_mask .unsqueeze (- 1 ),
728- )
729- target_d_global = torch .where (global_mask , d_target , torch .zeros_like (d_target ))
730- edge_inv_global = torch .where (global_mask , edge_inv_global , torch .zeros_like (edge_inv_global ))
731- target_pos_global = eq_transform (target_d_global , pos_perturbed , edge_index , edge_length )
732- node_eq_global = eq_transform (edge_inv_global , pos_perturbed , edge_index , edge_length )
733- loss_global = (node_eq_global - target_pos_global ) ** 2
734- loss_global = 2 * torch .sum (loss_global , dim = - 1 , keepdim = True )
735-
736- target_pos_local = eq_transform (
737- d_target [local_edge_mask ], pos_perturbed , edge_index [:, local_edge_mask ], edge_length [local_edge_mask ]
738- )
739- node_eq_local = eq_transform (
740- edge_inv_local , pos_perturbed , edge_index [:, local_edge_mask ], edge_length [local_edge_mask ]
741- )
742- loss_local = (node_eq_local - target_pos_local ) ** 2
743- loss_local = 5 * torch .sum (loss_local , dim = - 1 , keepdim = True )
744-
745- # loss for atomic eps regression
746- loss = loss_global + loss_local
747-
748- if return_unreduced_edge_loss :
749- pass
750- elif return_unreduced_loss :
751- return loss , loss_global , loss_local
752- else :
753- return loss
754-
755613 def get_residual_params (
756614 self ,
757615 t ,
0 commit comments