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¶
Pool edge features over the second axis into node features. |
|
Pool edge features (mean/min/max/std) into the global feature. |
|
Rescale coordinates by a learned function of their norm (SE(3)-safe). |
|
Normalize positions by their mean norm over real nodes. |
|
Pool node features (mean/min/max/std) into the global feature. |
Functions¶
|
Softmax over |
Module Contents¶
- class MolecularDiffusion.modules.models.midi.layers.EtoX(de: int, dx: int)¶
Bases:
torch.nn.ModulePool 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.ModulePool 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.ModuleRescale 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.ModuleNormalize positions by their mean norm over real nodes.
- forward(pos: torch.Tensor, node_mask: torch.Tensor) torch.Tensor¶
pos (B,N,3),node_mask (B,N,1)-> normalized positions.
- eps = 1e-05¶
- normalized_shape = (1,)¶
- weight¶
- class MolecularDiffusion.modules.models.midi.layers.Xtoy(dx: int, dy: int)¶
Bases:
torch.nn.ModulePool 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
xwithmask == 0entries driven to zero.