Skip to main content
The kumoai.trainer module provides the Trainer class for training custom GNN models on a Graph and TrainingTable, and generating predictions with a PredictionTable. Models can be customized with ModelPlan, though the plan suggested by PredictiveQuery.suggest_model_plan() is typically sufficient for strong out-of-the-box performance.

Model Plan

A ModelPlan defines the full parameter specification for training a Kumo model. It is composed of five sub-plans:
  • ColumnProcessingPlan — encoder overrides for individual columns
  • ModelArchitecturePlan — GNN or Graph Transformer parameters
  • NeighborSamplingPlan — subgraph sampling parameters
  • OptimizationPlan — learning rate, batch size, epochs, and related settings
  • TrainingJobPlan — AutoML-level settings
After generating a default model plan with PredictiveQuery.suggest_model_plan(), no further changes are required to train your first model. These options are available for fine-tuning.

ModelPlan

The top-level model configuration object. Each sub-plan is accessible as an attribute.
TrainingJobPlan
default:"TrainingJobPlan()"
AutoML job-level settings.
ColumnProcessingPlan
default:"ColumnProcessingPlan()"
Encoder overrides for individual columns.
NeighborSamplingPlan
default:"NeighborSamplingPlan()"
Subgraph sampling configuration.
OptimizationPlan
default:"OptimizationPlan()"
Optimization hyperparameters.
ModelArchitecturePlan
default:"ModelArchitecturePlan()"
GNN or Graph Transformer architecture parameters.

ColumnProcessingPlan

Specifies encoder overrides and missing value strategy overrides for individual table columns.
Optional[Dict[str, Encoder]]
default:"None"
A mapping from "table.column" to an encoder instance. Overrides Kumo’s auto-inferred encoder for that column.
Optional[Dict[Stype, NAStrategy]]
default:"None"
A mapping from semantic type to NAStrategy. Overrides the default imputation strategy for all columns of that semantic type.

ModelArchitecturePlan

Base class for architecture plans. Use GNNModelPlan or GraphTransformerModelPlan to configure a specific architecture.

GNNModelPlan

Configures a Graph Neural Network architecture.
List[int]
default:"inferred"
Candidate hidden channel sizes for AutoML search.
List[List[AggregationType]]
default:"inferred"
Candidate aggregation function combinations for AutoML search.
List[float]
default:"inferred"
Candidate dropout rates for AutoML search.

GraphTransformerModelPlan

Configures a Graph Transformer architecture.
List[int]
default:"inferred"
Candidate hidden channel sizes.
List[int]
default:"inferred"
Candidate number of transformer layers.
List[int]
default:"inferred"
Candidate number of attention heads.
List[float]
default:"inferred"
Candidate dropout rates.
List[List[PositionalEncodingType]]
default:"inferred"
Candidate positional encoding combinations.

NeighborSamplingPlan

Controls how Kumo samples subgraphs during training.
List[List[int]]
default:"inferred"
Candidate per-hop neighbor counts for AutoML search. Each inner list specifies the number of neighbors to sample at each hop.
bool
default:"inferred"
Whether to sample neighbors from the entity table.

OptimizationPlan

Controls learning rate, batch size, epochs, and other training optimization parameters.
int
default:"inferred"
Maximum number of training epochs.
int
default:"inferred"
Maximum number of training steps per epoch.
int
default:"inferred"
Maximum number of validation steps.
int
default:"inferred"
Maximum number of test steps.
List[Union[str, LossConfig]]
default:"inferred"
Candidate loss functions for AutoML search.
List[float]
default:"inferred"
Candidate base learning rates.
List[float]
default:"inferred"
Candidate weight decay values.
List[int]
default:"inferred"
Candidate batch sizes.
List[Optional[EarlyStoppingConfig]]
default:"inferred"
Candidate early stopping configurations.
List[Optional[LRSchedulerConfig]]
default:"inferred"
Candidate learning rate scheduler configurations.
List[Optional[float]]
default:"inferred"
Candidate majority class sampling ratios for imbalanced classification.
List[Optional[WeightMode]]
default:"inferred"
Candidate sample weighting modes.

TrainingJobPlan

AutoML job-level settings controlling the number of experiments and evaluation metrics.
int
default:"inferred"
Number of hyperparameter experiments to run during AutoML.
List[str]
default:"inferred"
Evaluation metrics to compute.
str
default:"inferred"
The primary metric used to select the best model.
bool
default:"True"
Whether to refit the best model on the combined train+validation set.
bool
default:"False"
Whether to additionally refit on the full dataset (train+validation+test).

Training

Trainer

Trains a Kumo GNN model on a PredictiveQuery. The two primary methods are fit() (training) and predict() (batch inference).
ModelPlan
required
The model plan specifying architecture, optimization, and sampling parameters.

model_plan property

Returns Optional[ModelPlan]

encoders property

Returns Optional[Dict[str, str]] — The encoder configuration used during training.

is_trained property

Returns boolTrue if this trainer has been successfully fit and is ready for prediction.

fit()

Trains a model on the provided graph and training table.
Graph
required
The relational graph.
Union[TrainingTable, TrainingTableJob]
default:"None"
An optional pre-generated training table. Provide exactly one of train_table or pquery.
Optional[PredictiveQuery]
default:"None"
An optional PredictiveQuery. When provided without train_table, Kumo generates the training table as an inline child job sharing the same graph snapshot. Recommended for live-streaming data to ensure consistency.
bool
default:"False"
If True, returns a TrainingJob immediately rather than blocking.
Mapping[str, str]
default:"{}"
Optional key-value tags attached to the training job.
str
default:"None"
Training job ID to warm-start from (initializes from an existing model’s weights).
Returns Union[TrainingJob, TrainingJobResult]

predict()

Generates batch predictions using the trained model.
Graph
required
The relational graph.
Union[PredictionTable, PredictionTableJob]
default:"None"
The prediction table generated from a PredictiveQuery.
Union[OutputConfig, Dict[str, Any]]
required
Output configuration for the batch prediction job (OutputConfig instance or dict). Key fields: output_connector (connector to write to), output_table_name (table name, or schema/table tuple for Databricks), output_types (predictions or embeddings), output_metadata_fields (JOB_TIMESTAMP or ANCHOR_TIMESTAMP).
Connector
default:"None"
Deprecated. Raises ValueError when passed. Use output_config instead.
str
default:"None"
Deprecated. Raises ValueError when passed. Use output_config instead.
bool
default:"False"
If True, returns a BatchPredictionJob immediately.
Mapping[str, str]
default:"{}"
Optional key-value tags attached to the prediction job.
str
default:"None"
The job ID of the training job whose model will be used for prediction. If None, uses the model from the most recent fit() call on this Trainer instance.
float
default:"None"
For binary classification models, the score threshold above which predictions are classified as 1. If None, the raw probability score is returned.
Optional[int]
default:"None"
For ranking task models, the number of classes to return in the prediction output.
int
default:"1"
Number of parallel workers for batch prediction. Values greater than 1 partition the prediction table and process in parallel.
datetime
default:"None"
The point in time at which to generate predictions. Only valid when prediction_table is None. If None, the anchor time is inferred from the latest available data. Cannot be specified together with prediction_table.
bool
default:"False"
If True, generates per-prediction feature attributions in addition to prediction outputs.
Returns Union[BatchPredictionJob, BatchPredictionJobResult]

load() classmethod

Loads a trained Trainer from a completed training job.
str
required
The training job ID.
Returns Trainer

TrainingJob

Represents an ongoing training job.

result()

Blocks until complete and returns the TrainingJobResult. Returns TrainingJobResult

status()

Returns JobStatusReport

cancel()

Cancels the running training job.

metrics_so_far()

Returns Optional[ModelEvaluationMetrics] — Metrics computed so far during training.

progress()

Returns AutoTrainerProgress — Detailed progress information.

TrainingJobResult

Represents a completed training job.
TrainingJobID
required
The training job ID.

id property

Returns TrainingJobID

model_plan property

Returns ModelPlan — The model plan used in this training job.

training_table property

Returns Union[TrainingTableJob, TrainingTable]

predictive_query property

Returns PredictiveQuery — The predictive query that defined this training job.

tracking_url property

Returns str — URL to the training job in the Kumo UI.

metrics()

Returns ModelEvaluationMetrics — Evaluation metrics for the completed job.

holdout_df()

Returns pd.DataFrame — The holdout dataset as a DataFrame.

explain()

Returns per-entity feature importances for a predictive query.
str
required
The PQL query string.
Sequence[Union[str, float, int]]
default:"None"
Entity indices to explain. Explains all entities if None.
RunMode
default:"RunMode.FAST"
The run mode for explanation computation.
List[int]
default:"None"
Per-hop neighbor counts for subgraph sampling.
Union[pd.Timestamp, Literal['entity']]
default:"None"
The anchor time for temporal explanation.
Returns pd.DataFrame

Batch Prediction

BatchPredictionJob

Represents an ongoing batch prediction job.

result()

Returns BatchPredictionJobResult

status()

Returns JobStatusReport

cancel()

Cancels the running batch prediction job.

BatchPredictionJobResult

Represents a completed batch prediction job.

data_df()

Returns pd.DataFrame — Prediction results.

data_urls()

Returns List[str] — Download URLs for prediction results.

summary()

Returns BatchPredictionJobSummary

Online Serving and Distillation

Distillation training and export_model() produce a serving bundle (online model directory and embeddings.parquet from batch prediction) in storage you control. Inference uses NVIDIA Triton Inference Server to load that bundle. See the Online Serving guide for the end-to-end flow.

DistillationTrainer

Trains a shallow model for online serving by reusing representations (embeddings) from a base GNN training job.
DistilledModelPlan
required
The distilled model plan.
str
required
The training job ID of the base GNN model to distill from.

is_trained property

Returns bool

fit()

Graph
required
The relational graph.
Union[TrainingTable, TrainingTableJob]
required
The training table.
bool
default:"False"
If True, returns a TrainingJob immediately.
Mapping[str, str]
default:"{}"
Optional job tags.
Returns Union[TrainingJob, TrainingJobResult]

load() classmethod

str
required
The training job ID.
Returns DistillationTrainer

DistilledModelPlan

Model plan for distillation. Composed of TrainingJobPlan, ColumnProcessingPlan, OptimizationPlan, DistillationPlan, and a distillation-specific architecture plan.

DistillationPlan

Configuration for the distillation process, specifying embedding keys, time offsets, and real-time interaction settings.
List[str]
default:"inferred"
Column keys used as embedding inputs.
TimeOffset
default:"inferred"
Maximum time offset for embedding lookups.
TimeOffset
default:"inferred"
Minimum time offset for embedding lookups.
Dict[str, int]
default:"{}"
Real-time interaction table configuration.

export_model()

Exports online serving model files and batch prediction embeddings to external storage for use with Triton Inference Server.
ModelOutputConfig
required
Specifies the training job, output path, and batch prediction job to bundle.
bool
default:"True"
If True, returns an ArtifactExportJob immediately.
Optional[Literal['pandas', 'mphf', 'lmdb']]
default:"None"
Override for the storage backend used for online-embedding artifacts. One of pandas, mphf, or lmdb. If omitted, uses the value set on config. mphf is the server-side default and recommended for large catalogs. pandas is simple but limited to catalogs below 100M IDs. lmdb has the lowest serving RAM footprint but higher cold-cache latency.
Optional[int]
default:"None"
Override for the wire format version of the exported online-serving model. This keyword is accepted and forwarded to the export config; it only takes effect where the deployment backend supports payload-version selection. When omitted, uses the value set on config.
Returns Union[ArtifactExportJob, ArtifactExportResult]

ModelOutputConfig

Output configuration for export_model(). Specifies output types, destination connector, and table name.
Set[str]
required
The output types to produce. Valid values: "predictions", "embeddings", or both.
Connector
default:"None"
The connector to write outputs to. Local download only if None.
Union[str, Tuple[str, str]]
default:"None"
Table name in the output connector. For Databricks, provide a (schema, table) tuple.
List[MetadataField]
default:"None"
Additional metadata columns to include in prediction output. Options: JOB_TIMESTAMP, ANCHOR_TIMESTAMP.

ArtifactExportJob

Represents an ongoing model artifact export job.

id property

Returns str

result()

Returns ArtifactExportResult

status()

Returns JobStatus

cancel()

Returns boolTrue if the job was successfully cancelled.

ArtifactExportResult

Represents a completed model artifact export.

tracking_url()

Returns str — URL to the export job in the Kumo UI.