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:
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:
| Field | Role |
|---|---|
arch | Architecture config (MLPConfig, …) |
embedding_config | Numerical / temporal embeddings |
loss, metric | Training loss and eval metric (symbols; see Losses) |
nrounds, early_stopping_rounds | Epoch budget and patience |
lr, wd, batchsize, seed | Optimiser and data |
scale_target | Standardize the target for losses that support it |
backend, device, gpuID | AD 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
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
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. AVector{Symbol}orVector{String}of the feature names to use.target_name: Required. ASymbolorStringindicating the name of the target variable.weight_name=nothing: Optional. ASymbolorStringindicating the sample weights column.offset_name=nothing: Optional. ASymbolorStringindicating 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 togroup_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 everyNepochs.verbosity=1: Integer. Controls the logging level (0for silent,>0for info).
Metric, early stopping, and device (device, gpuID) are taken from config, not from fit kwargs.
Inference
A fitted NeuroTabModel is callable. That is the usual path:
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
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 (:cpuor:gpu).backend: AD backend (:zygote,:enzyme, or:reactant). Defaults to the value stored on the model.proj=true: Whentrue, map raw outputs to natural scale. Set tofalsefor raw model-scale predictions.
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 anAbstractDataFrame.
Keyword arguments
device=:cpu: Execution device (:cpuor:gpu).backend: AD backend (:zygote,:enzyme, or:reactant). Defaults to the value stored on the model.proj=true: Whentrue, map raw outputs to natural scale. Set tofalsefor raw model-scale outputs.