MolecularDiffusion.modules.models.midi.layers

Dense-tensor building blocks for MiDi’s relational graph transformer.

These live here rather than in modules/layers/ on purpose: every one of them is hardcoded to MiDi’s dense (B,N,·) / (B,N,N,·) layout and is reused by nothing else in the platform.

Upstream’s SetNorm/GraphNorm are omitted – they are commented out at every call site in midi/models/transformer_model.py (lines 38-39, 50-51, 89, 97, 106, 111) in favour of plain LayerNorm, so they carry no weights in any released checkpoint.

Classes

EtoX

Pool edge features over the second axis into node features.

Etoy

Pool edge features (mean/min/max/std) into the global feature.

PositionsMLP

Rescale coordinates by a learned function of their norm (SE(3)-safe).

SE3Norm

Normalize positions by their mean norm over real nodes.

Xtoy

Pool node features (mean/min/max/std) into the global feature.

Functions

masked_softmax(→ torch.Tensor)

Softmax over x with mask == 0 entries driven to zero.

Module Contents

class MolecularDiffusion.modules.models.midi.layers.EtoX(de: int, dx: int)

Bases: torch.nn.Module

Pool edge features over the second axis into node features.

forward(E: torch.Tensor, e_mask2: torch.Tensor) torch.Tensor

E (B,N,N,de) -> (B,N,dx).

lin
class MolecularDiffusion.modules.models.midi.layers.Etoy(d: int, dy: int)

Bases: torch.nn.Module

Pool edge features (mean/min/max/std) into the global feature.

forward(E: torch.Tensor, e_mask1: torch.Tensor, e_mask2: torch.Tensor) torch.Tensor

E (B,N,N,de) -> (B,dy).

lin
class MolecularDiffusion.modules.models.midi.layers.PositionsMLP(hidden_dim: int, eps: float = 1e-05)

Bases: torch.nn.Module

Rescale coordinates by a learned function of their norm (SE(3)-safe).

forward(pos: torch.Tensor, node_mask: torch.Tensor) torch.Tensor

pos (B,N,3), node_mask (B,N) -> rescaled, re-centred pos.

eps = 1e-05
mlp
class MolecularDiffusion.modules.models.midi.layers.SE3Norm(eps: float = 1e-05, device: torch.device | None = None, dtype: torch.dtype | None = None)

Bases: torch.nn.Module

Normalize positions by their mean norm over real nodes.

extra_repr() str

Describe the layer for repr.

forward(pos: torch.Tensor, node_mask: torch.Tensor) torch.Tensor

pos (B,N,3), node_mask (B,N,1) -> normalized positions.

reset_parameters() None

Reset the single scale parameter to one.

eps = 1e-05
normalized_shape = (1,)
weight
class MolecularDiffusion.modules.models.midi.layers.Xtoy(dx: int, dy: int)

Bases: torch.nn.Module

Pool node features (mean/min/max/std) into the global feature.

forward(X: torch.Tensor, x_mask: torch.Tensor) torch.Tensor

X (B,N,dx), x_mask (B,N,1) -> (B,dy).

lin
MolecularDiffusion.modules.models.midi.layers.masked_softmax(x: torch.Tensor, mask: torch.Tensor, **kwargs: object) torch.Tensor

Softmax over x with mask == 0 entries driven to zero.