Skip to content

Layers ​

Reusable Lux layers shared by embeddings and backbones.

CarryMask, MaskSkip, MaskedBatchNorm, and AttnResidual take a mask that marks which columns are real observations versus padding. Grouped attention models use that so pad slots are ignored in attention and batch norm. See Grouped padding and masks.

NeuroTabModels.Models.Layers.AttnResidual Type
julia
AttnResidual(hsize, nheads; dropout=0.0, attn_dropout=0.0, attn_scale=0.1f0)

Peer attention over an unordered set (attention batch dim = 1). Shared for query and key, values = encoder tokens. Score scale is the usual NNlib.dot_product_attention . Residual x + scale * Dropout(Attn), with scale a learned scalar initialized at attn_scale (default 0.1). Encoder tokens should already be BatchNorm'd so QK logits stay O(1); the small residual scale is what lets the model ignore peers when the batch is not a real group.

Inputs are (hidden, seq) feature-first matrices. An optional key-padding mask may be passed as (x, mask) so padded group-buffer slots are ignored.

source
NeuroTabModels.Models.Layers.CarryMask Type
julia
CarryMask(layer)

Pass a padding mask through a layer that only acts on the feature matrix: (x, mask) ↦ (layer(x), mask). Unmasked x is forwarded unchanged.

source
NeuroTabModels.Models.Layers.GroupedDense Type
julia
GroupedDense(in_dims => out_dims, n_groups, activation=identity; kwargs...)

Independent dense map in_dims → out_dims for each of n_groups groups.

Shapes

  • Input: (in_dims, n_groups, batch)

  • Weight: (out_dims, in_dims, n_groups)

  • Bias: (out_dims, n_groups, 1)

  • Output: (out_dims, n_groups, batch)

Use this for packed MLP ensembles (n_groups = k) and for per-feature numerical embeddings (n_groups = nfeats). Shared-weight batch ensemble (LinearBatchEnsemble) is a different operator.

Keyword arguments

  • use_bias::Bool: Include bias (default true).

  • init_weight: Called as init_weight(rng, out_dims, in_dims, n_groups). Default rsqrt_uniform_grouped (fan-in of each group).

  • init_bias: Called as init_bias(rng, out_dims, n_groups, 1). Default zeros32.

source
NeuroTabModels.Models.Layers.MaskSkip Type
julia
MaskSkip(layer)

Skip-add that threads a padding mask: (x, mask) ↦ (x + f(x), mask).

source
NeuroTabModels.Models.Layers.MaskedBatchNorm Type
julia
MaskedBatchNorm(chs, act=identity; epsilon=1f-5, momentum=0.1f0)

BatchNorm over the last dimension of (channels, tokens).

When the input is (x, valid) with valid a per-token flag, mean/variance and running stats use only valid tokens so zero-padded group-buffer slots do not leak. Unmasked x is ordinary BatchNorm.

Dense / Dropout in the same chain should be wrapped in CarryMask; residual skips should use MaskSkip.

source
NeuroTabModels.Models.Layers.ResidualScale Type
julia
ResidualScale(init)

Learned scalar on a residual branch. One parameter, initialized at init (default 0.1) so peer attention starts as a small add-on rather than a full-scale mix of the batch / group.

source
NeuroTabModels.Models.Layers.glorot_uniform_grouped Method
julia
glorot_uniform_grouped(rng, out, in, groups)

Glorot/Xavier uniform using the 2D fan of each group (out × in), not conv nfan.

source
NeuroTabModels.Models.Layers.rsqrt_uniform_grouped Method
julia
rsqrt_uniform_grouped(rng, out, in, groups)

Independent uniform init on [-1/√in, 1/√in] for each group. Matches the previous NLinear / packed-TabM weight init; glorot_uniform on a 3D array would use conv-style nfan and fold groups into the fan.

source