Skip to content

MLP ​

NeuroTabModels.Models.MLP.MLPAttn Type
julia
MLPAttn

ResNet-style encoder (BatchNorm residual MLP) plus a thin peer-attention residual, then a linear head.

Scale: BatchNorm on the item stream (same recipe as ResNetConfig). n_attn_layers=0 is that encoder plus a Glorot head — the ResNet ablation.

Intended to sit after the usual embedding layer: Chain(embed, MLPAttn(...)).

Forward signatures

  • (x, ps, st) with x of shape (features, batch): all observations are valid tokens.

  • ((x, w), ps, st): w is a weight / boolean mask over the batch. Zero (or false) positions are treated as padded group-buffer slots and are ignored by attention via a rectangular key-padding mask of shape (seq, 1, 1, 1), broadcast onto attention scores (kv_len, q_len, nheads, 1). This is not a causal (triangular) mask. Encoder BatchNorm uses the same valid-token flags so pads do not enter mean/var (or running stats). The attention mask still zeros padded keys.

source
NeuroTabModels.Models.MLP.MLPAttnConfig Type
julia
MLPAttnConfig(; kwargs...)

Item-wise residual encoder with peer attention over the batch / group.

The encoder is a BatchNorm residual MLP like ResNetConfig, but each BN restricts mean/var to valid tokens when a padding mask is passed (grouped loaders). Attention is shared Q=K via NNlib.dot_product_attention, values = encoder tokens, residual-added with a learned scalar (attn_scale, default 0.1) so peer mixing starts small. The head is Glorot Dense like ResNet. n_attn_layers=0 should track ResNet on ungrouped data.

When a padding mask is available (w from grouped loaders, or the infer mask), the loss / eval / infer call sites pass (x, w) into the assembled MaskedModel.

Arguments

  • act::Symbol: Activation — :relu, :gelu, :sigmoid, or :tanh (default :relu).

  • hidden_size::Int: Encoder / attention dimension (default 64). Must be divisible by nheads.

  • stack_size::Int: Number of residual blocks after the stem (default 1). 0 is a no-op (embedding width must equal hidden_size).

  • dropout::Float64: Dropout in residual blocks and on the attention residual (default 0.0).

  • nheads::Int: Number of attention heads (default 4).

  • n_attn_layers::Int: Number of attention residuals (default 1). 0 is encoder + head only.

  • attn_dropout::Float64: Dropout on attention scores (default 0.0).

  • attn_scale::Float32: Initial value of the learned attention residual scale (default 0.1). Mixing is x + scale * Attn.

source
NeuroTabModels.Models.MLP.MLPAttnConfig Method
julia
(config::MLPAttnConfig)(; ins, outsize)

Build an MLPAttn backbone from config. fit prepends embeddings via MaskedModel(embed, core) when a padding mask must reach attention; otherwise Chain(embed, core) as for the other architectures.

source
NeuroTabModels.Models.MLP.MLPConfig Type
julia
MLPConfig(; kwargs...)

Configuration for a multi-layer perceptron backbone.

Arguments

  • act::Symbol: Activation — :relu, :gelu, :sigmoid, or :tanh (default :relu).

  • hidden_size::Int: Hidden dimension (default 64).

  • stack_size::Int: Number of hidden blocks (default 1).

  • dropout::Float64: Dropout rate between blocks (default 0.0).

  • MLE_tree_split::Bool: Split output head for Gaussian MLE (default false).

source
NeuroTabModels.Models.MLP.MLPConfig Method
julia
(config::MLPConfig)(; ins, outsize)

Build a Lux.Chain from config.

Arguments

  • ins::Int: Number of input features.

  • outsize::Int: Number of output units.

Returns

A Lux.Chain of Dense → BatchNorm → Dense blocks with optional dropout.

source
NeuroTabModels.Models.MLP._mlp_encoder Method
julia
_mlp_encoder(ins, hsize, act, stack_size, dropout)

ResNet trunk without the prediction Dense. Accepts x or (x, valid) so BatchNorm can skip padded group-buffer slots.

  • stack_size == 0: NoOpLayer. Requires ins == hsize.

  • stack_size >= 1: Dense(ins → hsize) + MaskedBN+act, then stack_size residual blocks.

source
NeuroTabModels.Models.MLP._res_block Method
julia
_res_block(hsize, act, dropout)

ResNet residual block with masked BatchNorm so grouped pads can be ignored.

source