NeuroTrees
NeuroTabModels.Models.NeuroTrees.MOETree Type
MOETreeMixture of NeuroTree experts with a tree router.
The router is one NeuroTree ensemble (k = 1) with outs = N logits. Softmax over that prediction axis yields a weight for each of the N experts. The experts are a classical NeuroTree with k = N independent ensembles. The mixture is a weighted sum over the ensemble axis, reducing k to 1.
Shapes (N experts, P outputs, batch B):
router:
(N, 1, B)→ softmax over dim 1experts:
(P, N, B)output:
(P, 1, B)
NeuroTabModels.Models.NeuroTrees.MOETreeConfig Type
MOETreeConfig(; kwargs...)Mixture of k NeuroTree experts gated by a softmax tree router.
Both branches see the same (embedded) features. The router is NeuroTree(ins => k; k = 1) — one ensemble whose leaf predictions are the k logits. Softmax over that axis produces the mixture weights. The experts are NeuroTree(ins => outsize; k = k) — k independent ensembles, mixed by those weights.
Arguments
tree_type::Symbol::binaryor:oblivious(default:binary).actA::Symbol: Feature activation on split weights. One of:identity,:tanh,:hardtanh, or:tanhshrink(default:identity).depth::Int: Tree depth (default4).ntrees::Int: Number of trees averaged in each ensemble (default32).k::Int: Number of experts (default4). Routeroutsand expert ensemble width. Must be ≥ 1.scaler::Bool: Apply softplus scaling on tree logits (defaulttrue).init_scale::Float32: Leaf weight init scale (default0.1).MLE_tree_split::Bool: Split output head for Gaussian MLE (defaultfalse).
NeuroTabModels.Models.NeuroTrees.MOETreeConfig Method
(config::MOETreeConfig)(; ins, outsize)Build a MOETree backbone from config.
NeuroTabModels.Models.NeuroTrees.NeuroTree Type
NeuroTree(feats => outs; tree_type=:binary, actA=identity, scaler=true,
depth, trees, k, init_scale=0.1)Differentiable tree ensemble layer. Output dims: [outs, k, batch_size].
Arguments
feats::Int: Number of input features.outs::Int: Number of predictions per leaf (thePaxis). Encoder use isouts = 1.tree_type::Symbol::binaryor:oblivious.actA: Feature activation applied to split weights.scaler::Bool: Scale logits with a learned softplus factor.depth::Int: Tree depth.trees::Int: Number of trees averaged in each of thekensembles.k::Int: Number of independent ensembles. Each ensemble produces oneouts-wide vector; leaf values are not shared acrossk.init_scale::Float32: Standard deviation for leaf weight initialization (default0.1).
NeuroTabModels.Models.NeuroTrees.NeuroTreeAttn Type
NeuroTreeAttnNeuroTree encoder (per-observation numerical embeddings) followed by a per-channel peer-attention residual, then a prediction head.
Intended to sit after the usual embedding layer: Chain(embed, NeuroTreeAttn(...)). Same role as MLPAttn: a per-row map into hidden_size, peer attention over the batch / group. For a scalar head the k ensembles are kept through the loss, same layout as NeuroTreeConfig.
The encoder width is NeuroTree's k axis, not outs: NeuroTree(ins => 1; k = hidden_size). Each hidden channel is its own ensemble (independent splits and leaf values). Putting hidden_size on outs with k = 1 would share one routing among all channels. Native layout is (1, k, batch); FlattenLayer only drops the singleton outs axis so attention sees (hidden_size, batch). After attention, the k axis is restored to (1, k, batch) so the loss trains each channel as an independent predictor, same as NeuroTreeConfig. A Dense+BatchNorm head would collapse that axis.
Those tokens are the sequence of a single batch / group, identical to MLPAttn.
Forward signatures
(x, ps, st)withxof shape(features, batch): all observations are valid tokens.((x, w), ps, st):wis a weight / boolean mask over the batch. Zero (orfalse) positions are treated as padded group-buffer slots and are ignored by attention via a rectangular key-padding mask.
NeuroTabModels.Models.NeuroTrees.NeuroTreeAttnConfig Type
NeuroTreeAttnConfig(; kwargs...)Configuration for a NeuroTree encoder plus batch-level transformer attention.
The tree stem is NeuroTree(ins => 1; k = hidden_size): hidden_size independent ensembles (own splits per channel). FlattenLayer reshapes (1, k, batch) → (k, batch) for attention, then the head restores (1, k, batch) so MSE trains each ensemble against y (inference means over k). Peer attention is per ensemble channel (one head per k, no Dense QK mixing), residual-added with attn_scale (default 0). There is no transformer FFN.
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
tree_type::Symbol::binaryor:oblivious(default:binary).actA::Symbol: Feature activation on split weights. One of:identity,:tanh,:hardtanh, or:tanhshrink(default:identity).depth::Int: Tree depth (default4). Controls the number of leaves (2^depth), which is an internal routing axis — not the hidden width.ntrees::Int: Number of trees averaged in each of thekhidden ensembles (default32).hidden_size::Int: Encoding / attention dimension (default64). Equals encoder NeuroTreek(outs = 1), so each channel has its own splits. Peer attention uses one head per channel (nheadsis ignored).stack_size::Int: Encoder depth (default1).0is a no-op (embedding width must equalhidden_size).1is a singleNeuroTree+ flatten. Each extra layer is a residualNeuroTreeof widthhidden_size, with optional dropout.scaler::Bool: Apply softplus scaling on tree logits (defaulttrue).init_scale::Float32: Leaf weight init scale (default0.1).dropout::Float64: Dropout after extra encoder layers only (stack_size ≥ 2). Not applied to the attention residual (those tokens are thekpredictions).nheads::Int: Ignored. Peer attention uses one head per ensemble channel.n_attn_layers::Int: Number of attention residuals (default1).0skips attention.attn_dropout::Float64: Dropout on attention scores (default0.0).attn_scale::Float32: Initial value of the learned attention residual scale (default0.0). Mixing isx + scale * Attnwith one head per ensemble channel (confignheadsis not used). Default0son_attn_layers=1starts asn_attn_layers=0. Raise it if grouped peers should mix.
NeuroTabModels.Models.NeuroTrees.NeuroTreeAttnConfig Method
(config::NeuroTreeAttnConfig)(; ins, outsize)Build a NeuroTreeAttn 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.
NeuroTabModels.Models.NeuroTrees.NeuroTreeConfig Type
NeuroTreeConfig(; kwargs...)Configuration for differentiable neuro-tree ensembles.
Arguments
tree_type::Symbol::binaryor:oblivious(default:binary).actA::Symbol: Feature activation. One of:identity,:tanh,:hardtanh, or:tanhshrink(default:identity).depth::Int: Tree depth (default4).ntrees::Int: Number of trees per layer (default32).k::Int: Ensemble size.hidden_size::Int: Hidden dimension for stacked trees (default1).stack_size::Int: Number of stacked tree layers (default1).scaler::Bool: Apply softplus scaling on tree logits (defaulttrue).init_scale::Float32: Leaf weight init scale (default0.1).MLE_tree_split::Bool: Split output head for Gaussian MLE (defaultfalse).
NeuroTabModels.Models.NeuroTrees.NeuroTreeConfig Method
(config::NeuroTreeConfig)(; 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 stacked neuro-tree layers.
NeuroTabModels.Models.NeuroTrees._tree_attn_block Method
_tree_attn_block(ins, hsize, tree_kwargs)One NeuroTree encoder block: k = hsize independent ensembles, outs = 1 (scalar per leaf, own splits per channel). FlattenLayer only drops the singleton outs axis ((1, k, batch) → (k, batch)) so attention sees a 2D token matrix. Wrapped in CarryMask so a padding flag still reaches attention after the encoder.
NeuroTabModels.Models.NeuroTrees._tree_attn_encoder Method
_tree_attn_encoder(ins, hsize, stack_size, dropout, tree_kwargs)Per-observation map into the attention width hsize.
Each NeuroTree produces (1, k, batch) with k = hsize. FlattenLayer reshapes to (hsize, batch). Same adapter as stacked NeuroTree hidden layers.
stack_size == 0:NoOpLayer. Requiresins == hsizeso the NeuroTab embedding block can be the sole numerical embedding.stack_size == 1:_tree_attn_block(ins, hsize). No encoder dropout.stack_size >= 2: that stem, thenstack_size - 1residualNeuroTreeblocks of widthhsize, with optional dropout after each residual.
NeuroTabModels.Models.NeuroTrees.get_logits_mask Method
get_logits_mask(::Val{:binary}, depth::Integer)NeuroTabModels.Models.NeuroTrees.get_softplus_mask Method
get_softplus_mask(::Val{:binary}, depth::Integer)