Skip to content

Commit 0bbedfe

Browse files
committed
remove unused code
1 parent 1088818 commit 0bbedfe

1 file changed

Lines changed: 2 additions & 144 deletions

File tree

src/diffusers/models/dualencoder_gfn.py

Lines changed: 2 additions & 144 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,6 @@
1313
from torch_geometric.utils import dense_to_sparse, to_dense_adj
1414
from torch_scatter import scatter_add, scatter_mean
1515
from torch_sparse import SparseTensor, coalesce
16-
from tqdm.auto import tqdm
1716

1817
from ..configuration_utils import ConfigMixin
1918
from ..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-
492462
class 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

Comments
 (0)