DeepMTP.models package

The canonical model API groups reusable branch networks, combination architectures, and construction:

from DeepMTP.models import (
    BranchEncoder,
    CompositeEncoder,
    ConvNet,
    FUSION_REGISTRY,
    GraphEncoder,
    MLP,
    ModelFactory,
    SparseMLP,
    TabularEncoder,
)

Encoder contracts

All built-in branch encoders declare an input kind and output dimension. Their shared runtime contract requires a floating output shaped [batch_size, output_dim].

Most encoders consume one tensor. SparseMLP consumes a PyTorch COO or CSR sparse tensor and produces a dense representation after its first projection. GraphEncoder consumes a structured GraphInput and applies GIN or GINE message passing followed by graph-level pooling. TabularEncoder consumes a structured TabularInput containing a floating numeric tensor and an integer categorical tensor while producing the same standardized encoder output. CompositeEncoder consumes a named CompositeInput, validates every child encoder output, and concatenates the representations in declared order.

Custom branch factories should declare input_kind and output_dim on their returned module. For compatibility, the model factory can infer a missing non-dot-product output_dim from a final Linear layer, but this structural inference is deprecated.

Shared contracts for branch encoders and their standardized outputs.

class DeepMTP.models.contracts.BranchEncoder(*args, **kwargs)

Bases: Protocol

Structural contract implemented by built-in branch encoders.

input_kind: Literal['dense', 'id', 'image', 'graph', 'sequence', 'sparse', 'tabular', 'custom', 'composite']
output_dim: int
DeepMTP.models.contracts.encode_branch(encoder: torch.nn.Module, values: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput, *, branch: str, explicit_dimension: int | None = None) → torch.Tensor

Run one encoder and validate the common [batch, width] contract.

DeepMTP.models.contracts.encoder_input_kind(encoder: torch.nn.Module, *, branch: str) → Literal['dense', 'id', 'image', 'graph', 'sequence', 'sparse', 'tabular', 'custom', 'composite']

Return a validated declared encoder input modality.

DeepMTP.models.contracts.encoder_output_dimension(encoder: torch.nn.Module, *, branch: str, explicit_dimension: int | None = None) → int

Return a validated declared encoder output dimension.

Branch models

class DeepMTP.models.branches.CompositeEncoder(*args: Any, **kwargs: Any)

Bases: Module

Fuse standardized outputs from named component encoders.

forward(values: CompositeInput) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'composite'
class DeepMTP.models.branches.ConvNet(*args: Any, **kwargs: Any)

Bases: Sequential

A convolutional neural network that is based on resnet.

forward(v: torch.Tensor) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'image'
class DeepMTP.models.branches.GraphEncoder(*args: Any, **kwargs: Any)

Bases: Module

Encode homogeneous graphs with GIN/GINE message passing and pooling.

forward(values: GraphInput) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'graph'
class DeepMTP.models.branches.IDEmbedding(*args: Any, **kwargs: Any)

Bases: Module

Map zero-based entity IDs to trainable dense representations.

forward(entity_ids: torch.Tensor) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'id'
class DeepMTP.models.branches.MLP(*args: Any, **kwargs: Any)

Bases: Sequential

A standard fully connected feed-forward neural network.

forward(v: torch.Tensor) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'dense'
class DeepMTP.models.branches.SequenceConv1DEncoder(*args: Any, **kwargs: Any)

Bases: Module

Encode padded token sequences with masked temporal convolutions.

forward(values: SequenceInput) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'sequence'
class DeepMTP.models.branches.SequenceGRUEncoder(*args: Any, **kwargs: Any)

Bases: Module

Encode padded token sequences with an embedding layer and GRU.

forward(values: SequenceInput) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'sequence'
class DeepMTP.models.branches.SequenceTransformerEncoder(*args: Any, **kwargs: Any)

Bases: Module

Encode padded token sequences with masked self-attention.

forward(values: SequenceInput) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'sequence'
class DeepMTP.models.branches.SparseMLP(*args: Any, **kwargs: Any)

Bases: Module

Project sparse inputs before applying a dense MLP tail.

forward(values: torch.Tensor) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'sparse'
class DeepMTP.models.branches.TabularEncoder(*args: Any, **kwargs: Any)

Bases: Module

Encode normalized numeric values and embedded categorical columns.

forward(values: TabularInput) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'tabular'

Combination models and construction

FUSION_REGISTRY resolves the built-in dot_product, mlp, and kronecker combination models. ModelFactory uses this registry instead of maintaining a separate architecture conditional.

Two-branch model definitions and construction.

class DeepMTP.models.factory.ModelBundle(instance_branch: torch.nn.Module, target_branch: torch.nn.Module, model: torch.nn.Module)

Bases: object

The two branches and their combined model.

instance_branch: torch.nn.Module
model: torch.nn.Module
target_branch: torch.nn.Module
class DeepMTP.models.factory.ModelFactory(config: DeepMTPConfig | Mapping[str, Any])

Bases: object

Construct a complete DeepMTP model from validated configuration.

build(*, instance_branch_factory: Callable[[Mapping[str, Any]], torch.nn.Module] | None = None, target_branch_factory: Callable[[Mapping[str, Any]], torch.nn.Module] | None = None) → ModelBundle

Build branches and their configured combination model.

class DeepMTP.models.factory.TwoBranchDotProductModel(*args: Any, **kwargs: Any)

Bases: Module

Combine equal-sized branch embeddings with a dot product.

forward(instance_features: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput, target_features: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput) → torch.Tensor
class DeepMTP.models.factory.TwoBranchKroneckerModel(*args: Any, **kwargs: Any)

Bases: Module

Combine branch embeddings using batched Kronecker products.

forward(instance_features: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput, target_features: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput) → torch.Tensor
class DeepMTP.models.factory.TwoBranchMLPModel(*args: Any, **kwargs: Any)

Bases: Module

Combine branch embeddings using a multilayer perceptron.

forward(instance_features: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput, target_features: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput) → torch.Tensor
DeepMTP.models.factory.isolated_torch_seed(seed: int | None) → Iterator[None]

Seed model initialization without changing process-wide RNG state.

Package contents

Branch models, combination architectures, and model construction.

class DeepMTP.models.BranchEncoder(*args, **kwargs)

Bases: Protocol

Structural contract implemented by built-in branch encoders.

input_kind: Literal['dense', 'id', 'image', 'graph', 'sequence', 'sparse', 'tabular', 'custom', 'composite']
output_dim: int
class DeepMTP.models.CompositeEncoder(*args: Any, **kwargs: Any)

Bases: Module

Fuse standardized outputs from named component encoders.

forward(values: CompositeInput) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'composite'
class DeepMTP.models.ConvNet(*args: Any, **kwargs: Any)

Bases: Sequential

A convolutional neural network that is based on resnet.

forward(v: torch.Tensor) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'image'
class DeepMTP.models.GraphEncoder(*args: Any, **kwargs: Any)

Bases: Module

Encode homogeneous graphs with GIN/GINE message passing and pooling.

forward(values: GraphInput) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'graph'
class DeepMTP.models.IDEmbedding(*args: Any, **kwargs: Any)

Bases: Module

Map zero-based entity IDs to trainable dense representations.

forward(entity_ids: torch.Tensor) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'id'
class DeepMTP.models.MLP(*args: Any, **kwargs: Any)

Bases: Sequential

A standard fully connected feed-forward neural network.

forward(v: torch.Tensor) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'dense'
class DeepMTP.models.ModelBundle(instance_branch: torch.nn.Module, target_branch: torch.nn.Module, model: torch.nn.Module)

Bases: object

The two branches and their combined model.

instance_branch: torch.nn.Module
model: torch.nn.Module
target_branch: torch.nn.Module
class DeepMTP.models.ModelFactory(config: DeepMTPConfig | Mapping[str, Any])

Bases: object

Construct a complete DeepMTP model from validated configuration.

build(*, instance_branch_factory: Callable[[Mapping[str, Any]], torch.nn.Module] | None = None, target_branch_factory: Callable[[Mapping[str, Any]], torch.nn.Module] | None = None) → ModelBundle

Build branches and their configured combination model.

class DeepMTP.models.SequenceConv1DEncoder(*args: Any, **kwargs: Any)

Bases: Module

Encode padded token sequences with masked temporal convolutions.

forward(values: SequenceInput) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'sequence'
class DeepMTP.models.SequenceGRUEncoder(*args: Any, **kwargs: Any)

Bases: Module

Encode padded token sequences with an embedding layer and GRU.

forward(values: SequenceInput) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'sequence'
class DeepMTP.models.SequenceTransformerEncoder(*args: Any, **kwargs: Any)

Bases: Module

Encode padded token sequences with masked self-attention.

forward(values: SequenceInput) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'sequence'
class DeepMTP.models.SparseMLP(*args: Any, **kwargs: Any)

Bases: Module

Project sparse inputs before applying a dense MLP tail.

forward(values: torch.Tensor) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'sparse'
class DeepMTP.models.TabularEncoder(*args: Any, **kwargs: Any)

Bases: Module

Encode normalized numeric values and embedded categorical columns.

forward(values: TabularInput) → torch.Tensor
input_kind: ClassVar[BranchInputKind] = 'tabular'
class DeepMTP.models.TwoBranchDotProductModel(*args: Any, **kwargs: Any)

Bases: Module

Combine equal-sized branch embeddings with a dot product.

forward(instance_features: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput, target_features: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput) → torch.Tensor
class DeepMTP.models.TwoBranchKroneckerModel(*args: Any, **kwargs: Any)

Bases: Module

Combine branch embeddings using batched Kronecker products.

forward(instance_features: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput, target_features: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput) → torch.Tensor
class DeepMTP.models.TwoBranchMLPModel(*args: Any, **kwargs: Any)

Bases: Module

Combine branch embeddings using a multilayer perceptron.

forward(instance_features: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput, target_features: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput) → torch.Tensor
DeepMTP.models.encode_branch(encoder: torch.nn.Module, values: torch.Tensor | CompositeInput | GraphInput | SequenceInput | TabularInput, *, branch: str, explicit_dimension: int | None = None) → torch.Tensor

Run one encoder and validate the common [batch, width] contract.

DeepMTP.models.encoder_input_kind(encoder: torch.nn.Module, *, branch: str) → Literal['dense', 'id', 'image', 'graph', 'sequence', 'sparse', 'tabular', 'custom', 'composite']

Return a validated declared encoder input modality.

DeepMTP.models.encoder_output_dimension(encoder: torch.nn.Module, *, branch: str, explicit_dimension: int | None = None) → int

Return a validated declared encoder output dimension.

DeepMTP.models.isolated_torch_seed(seed: int | None) → Iterator[None]

Seed model initialization without changing process-wide RNG state.