Skip to content

Training ​

Training is a learner plus a DataFrame. The learner holds the architecture config and the training settings. fit builds the Lux chain, runs the loop, and returns a NeuroTabModel. Predictions come from calling that fitted model, or from infer.

Architecture configs and how they become a Lux chain are on Models. Losses are on Losses. Numerical and temporal embeddings are on Embeddings.

Learners ​

NeuroTabRegressor and NeuroTabClassifier are the constructors. They take an architecture config and the knobs that stay fixed for the run:

julia
using NeuroTabModels, NeuroTabModels.Models

cfg = MLPConfig(; hidden_size=64, stack_size=2)
learner = NeuroTabRegressor(
    cfg;
    loss=:mse,
    nrounds=100,
    lr=1.0f-2,
    batchsize=2048,
    backend=:zygote,
    device=:cpu,
)

These live on the learner, not on fit:

FieldRole
archArchitecture config (MLPConfig, …)
embedding_configNumerical / temporal embeddings
loss, metricTraining loss and eval metric (symbols; see Losses)
nrounds, early_stopping_roundsEpoch budget and patience
lr, wd, batchsize, seedOptimiser and data
scale_targetStandardize the target for losses that support it
backend, device, gpuIDAD backend (:zygote, :enzyme, :reactant) and device

metric and early_stopping_rounds only take effect when fit is given deval. Constructors and the MLJ interface are documented with the types on Models.

Fit ​

julia
m = NeuroTabModels.fit(
    learner,
    dtrain;
    feature_names,
    target_name,
    deval=deval,                 # optional: metrics and early stopping
    weight_name=nothing,
    offset_name=nothing,
    group_name=nothing,          # grouped / padded loader
    eval_group_name=group_name,
    print_every_n=10,
)

feature_names and target_name are required. dtrain (and deval) must be <:AbstractDataFrame.

group_name switches the dataloader to one padded group per step. That is required for attention models that mix across observations in a group; see Grouped padding and masks. eval_group_name can differ from group_name when you want grouped eval metrics while training on the ungrouped loader.

MLJModelInterface.fit Function
julia
fit(
    config::LearnerTypes,
    dtrain;
    feature_names,
    target_name,
    weight_name=nothing,
    offset_name=nothing,
    group_name=nothing,
    eval_group_name=group_name,
    deval=nothing,
    print_every_n=9999,
    verbosity=1,
)

Training function of NeuroTabModels' internal API.

Arguments

  • config::LearnerTypes: The configuration object defining the model architecture, loss, and training hyperparameters.

  • dtrain: The training data. Must be <:AbstractDataFrame.

Keyword arguments

  • feature_names: Required. A Vector{Symbol} or Vector{String} of the feature names to use.

  • target_name: Required. A Symbol or String indicating the name of the target variable.

  • weight_name=nothing: Optional. A Symbol or String indicating the sample weights column.

  • offset_name=nothing: Optional. A Symbol or String indicating the offset column.

  • group_name=nothing: Optional. Column used to group training data in the dataloader.

  • eval_group_name=group_name: Optional. Column used to group evaluation data when computing metrics. Defaults to group_name. Set independently to compute groupby eval metrics while training on the regular (ungrouped) dataloader.

  • deval=nothing: Optional. Evaluation data (<:AbstractDataFrame) for tracking metrics and early stopping.

  • print_every_n=9999: Integer. Logs training progress to the console every N epochs.

  • verbosity=1: Integer. Controls the logging level (0 for silent, >0 for info).

Metric, early stopping, and device (device, gpuID) are taken from config, not from fit kwargs.

source

Inference ​

A fitted NeuroTabModel is callable. That is the usual path:

julia
p = m(dtrain)                    # natural-scale predictions
p_raw = m(dtrain; proj=false)    # model-scale (logits, log-σ, …)

infer is the same function, also overloaded on an iterable of feature batches. proj=true (default) applies the inverse link and, when training used scale_target, undoes target scaling. device / backend default to the values stored on the model.

If the model was trained with group_name, inference groups the DataFrame the same way and drops pad slots after the forward.

NeuroTabModels.Infer.infer Function
julia
infer(m::NeuroTabModel, data; device=:cpu, backend=get(m.info, :backend, :zygote), proj=true)

Run inference on batched feature data.

Arguments

  • m::NeuroTabModel: A fitted model.

  • data: Iterable of feature batches.

Keyword arguments

  • device=:cpu: Execution device (:cpu or :gpu).

  • backend: AD backend (:zygote, :enzyme, or :reactant). Defaults to the value stored on the model.

  • proj=true: When true, map raw outputs to natural scale. Set to false for raw model-scale predictions.

source
julia
infer(m::NeuroTabModel, df::AbstractDataFrame; device=:cpu, backend=get(m.info, :backend, :zygote), proj=true)

Run inference on tabular data and return predictions.

Arguments

  • m::NeuroTabModel: A fitted model.

  • df: Feature data as an AbstractDataFrame.

Keyword arguments

  • device=:cpu: Execution device (:cpu or :gpu).

  • backend: AD backend (:zygote, :enzyme, or :reactant). Defaults to the value stored on the model.

  • proj=true: When true, map raw outputs to natural scale. Set to false for raw model-scale outputs.

source