MolecularDiffusion.modules.models.syncogen.diffusion.sampling.discrete_strategies.base¶
Base class for discrete sampling strategies.
Classes¶
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.ABCBase 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.
- 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¶