MolecularDiffusion.modules.models.syncogen.diffusion.sampling.discrete_strategies.base

Base class for discrete sampling strategies.

Classes

DiscreteStrategyBase

Base class for discrete (graph) sampling strategies.

Module Contents

class MolecularDiffusion.modules.models.syncogen.diffusion.sampling.discrete_strategies.base.DiscreteStrategyBase(discrete_noise: Callable | None = None, constrain_edge_sampling: bool = False)

Bases: abc.ABC

Base class for discrete (graph) sampling strategies.

Samplers accept BBRxnGraph objects and extract what they need internally. This keeps the call site clean and allows different strategies to access different properties (masks, indices, etc.) as needed.

Parameters:
  • discrete_noise – Noise schedule function for discrete features.

  • constrain_edge_sampling – Whether to constrain edge sampling step.

static edge_mask_index(D_edge: int) int
static edge_no_edge_index(D_edge: int) int
static node_mask_index(D_node: int) int
abstractmethod step(graph: MolecularDiffusion.modules.models.syncogen.api.graph.graph.BBRxnGraph, p_x0: torch.Tensor, p_e0: torch.Tensor, t: torch.Tensor, dt: torch.Tensor = None) Tuple[torch.Tensor, torch.Tensor]

Perform one denoising step for discrete features. :param graph: BBRxnGraph containing current noisy discrete features :param p_x0: Predicted node probabilities (usually softmax/logits.exp) (B, N, D_node) :param p_e0: Predicted edge probabilities (B, N, N, D_edge) :param t: Current timestep (B, 1) :param dt: Time step size (optional, some strategies may require it)

Returns:

Updated node features (B, N, D_node) E_next: Updated edge features (B, N, N, D_edge)

Return type:

X_next

constrain_edge_sampling = False
discrete_noise = None