Skip to content

NeuroTrees ​

NeuroTabModels.Models.NeuroTrees.MOETree Type
julia
MOETree

Mixture 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 1

  • experts: (P, N, B)

  • output: (P, 1, B)

source
NeuroTabModels.Models.NeuroTrees.MOETreeConfig Type
julia
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: :binary or :oblivious (default :binary).

  • actA::Symbol: Feature activation on split weights. One of :identity, :tanh, :hardtanh, or :tanhshrink (default :identity).

  • depth::Int: Tree depth (default 4).

  • ntrees::Int: Number of trees averaged in each ensemble (default 32).

  • k::Int: Number of experts (default 4). Router outs and expert ensemble width. Must be ≥ 1.

  • scaler::Bool: Apply softplus scaling on tree logits (default true).

  • init_scale::Float32: Leaf weight init scale (default 0.1).

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

source
NeuroTabModels.Models.NeuroTrees.MOETreeConfig Method
julia
(config::MOETreeConfig)(; ins, outsize)

Build a MOETree backbone from config.

source
NeuroTabModels.Models.NeuroTrees.NeuroTree Type
julia
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 (the P axis). Encoder use is outs = 1.

  • tree_type::Symbol: :binary or :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 the k ensembles.

  • k::Int: Number of independent ensembles. Each ensemble produces one outs-wide vector; leaf values are not shared across k.

  • init_scale::Float32: Standard deviation for leaf weight initialization (default 0.1).

source
NeuroTabModels.Models.NeuroTrees.NeuroTreeAttn Type
julia
NeuroTreeAttn

NeuroTree 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) 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.

source
NeuroTabModels.Models.NeuroTrees.NeuroTreeAttnConfig Type
julia
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: :binary or :oblivious (default :binary).

  • actA::Symbol: Feature activation on split weights. One of :identity, :tanh, :hardtanh, or :tanhshrink (default :identity).

  • depth::Int: Tree depth (default 4). 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 the k hidden ensembles (default 32).

  • hidden_size::Int: Encoding / attention dimension (default 64). Equals encoder NeuroTree k (outs = 1), so each channel has its own splits. Peer attention uses one head per channel (nheads is ignored).

  • stack_size::Int: Encoder depth (default 1). 0 is a no-op (embedding width must equal hidden_size). 1 is a single NeuroTree + flatten. Each extra layer is a residual NeuroTree of width hidden_size, with optional dropout.

  • scaler::Bool: Apply softplus scaling on tree logits (default true).

  • init_scale::Float32: Leaf weight init scale (default 0.1).

  • dropout::Float64: Dropout after extra encoder layers only (stack_size ≥ 2). Not applied to the attention residual (those tokens are the k predictions).

  • nheads::Int: Ignored. Peer attention uses one head per ensemble channel.

  • n_attn_layers::Int: Number of attention residuals (default 1). 0 skips attention.

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

  • attn_scale::Float32: Initial value of the learned attention residual scale (default 0.0). Mixing is x + scale * Attn with one head per ensemble channel (config nheads is not used). Default 0 so n_attn_layers=1 starts as n_attn_layers=0. Raise it if grouped peers should mix.

source
NeuroTabModels.Models.NeuroTrees.NeuroTreeAttnConfig Method
julia
(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.

source
NeuroTabModels.Models.NeuroTrees.NeuroTreeConfig Type
julia
NeuroTreeConfig(; kwargs...)

Configuration for differentiable neuro-tree ensembles.

Arguments

  • tree_type::Symbol: :binary or :oblivious (default :binary).

  • actA::Symbol: Feature activation. One of :identity, :tanh, :hardtanh, or :tanhshrink (default :identity).

  • depth::Int: Tree depth (default 4).

  • ntrees::Int: Number of trees per layer (default 32).

  • k::Int: Ensemble size.

  • hidden_size::Int: Hidden dimension for stacked trees (default 1).

  • stack_size::Int: Number of stacked tree layers (default 1).

  • scaler::Bool: Apply softplus scaling on tree logits (default true).

  • init_scale::Float32: Leaf weight init scale (default 0.1).

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

source
NeuroTabModels.Models.NeuroTrees.NeuroTreeConfig Method
julia
(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.

source
NeuroTabModels.Models.NeuroTrees._tree_attn_block Method
julia
_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.

source
NeuroTabModels.Models.NeuroTrees._tree_attn_encoder Method
julia
_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. Requires ins == hsize so 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, then stack_size - 1 residual NeuroTree blocks of width hsize, with optional dropout after each residual.

source
NeuroTabModels.Models.NeuroTrees.get_logits_mask Method
julia
get_logits_mask(::Val{:binary}, depth::Integer)
source
NeuroTabModels.Models.NeuroTrees.get_softplus_mask Method
julia
get_softplus_mask(::Val{:binary}, depth::Integer)
source