Skip to main content

Predict

talos.Predict selects a trained model from Scan or a RunResult and performs inference on caller-supplied features. Import it from talos. Construction stores the run; prediction happens when .predict() or .predict_classes() is called.

Prerequisites are a completed run with a recoverable model and its installed framework backend. Apply the same preprocessing and feature structure used for training; this wrapper does not reconstruct caller-owned preprocessing automatically.

Examples below use the held-out Iris setup in Scan → Minimal Example. Run that setup first; it defines scan_object, p, input_model, x, y, x_val, y_val, x_test, and y_test.

predictor = talos.Predict(scan_object)
probabilities = predictor.predict(x_test, metric='val_loss', asc=True)

Prediction methods​

predict returns the trained model's raw outputs for x. These are probabilities only when the model produces probabilities; regression predictions and logits remain unchanged.

predictor.predict(x_test, metric='val_loss', asc=True)

predict_classes makes class predictions on x which has to be in the same form as the input data used in the Scan() experiment.

predictor.predict_classes(x_test, metric='val_loss', asc=True, task='multi_class')

Arguments​

ParameterDefaultDescription
xrequiredFeatures in the model's expected input structure.
metricrequiredScan result column used to select a model.
ascrequiredTrue for lower metric values, False for higher values.
model_idNoneExplicit result-row index; otherwise select by metric.
taskrequired for predict_classesClass conversion rule described below.
savedFalseRecover from persisted artifacts rather than a retained live model.
custom_objectsNoneKeras objects needed for reconstruction.
model_factoryNoneTorch reconstruction factory when needed.
**kwargsempty.predict() only: forwarded to backend prediction, such as Keras batch_size.

Signatures and class conversion​

The constructor is Predict(scan_object). Its methods are predict(x, metric, asc, model_id=None, saved=False, custom_objects=None, model_factory=None, **kwargs) and predict_classes(x, metric, asc, task, model_id=None, saved=False, custom_objects=None, model_factory=None).

TaskClass conversion
binaryargmax for two outputs; otherwise threshold each value at 0.5.
multi_class, multiclass, multi_labelargmax across the last axis for multidimensional output; one-dimensional output becomes integers.
multilabel, multi_label_independentIndependent 0.5 thresholds.
continuous, regressionReturn the raw predicted values.

The historical multi_label spelling means exclusive class conversion here. Use multilabel or multi_label_independent for independent labels. Thresholds require suitably scaled outputs; Torch logits are not transformed into probabilities automatically. Structured multi-output predictions can be obtained through .predict(); class conversion expects a single array.

Prediction returns outputs and does not add columns to the result table or retrain the model. Tied selection metrics preserve trial order, and missing values are excluded. Missing metric columns raise KeyError; an unusable selection or unretained model raises ValueError. Unknown class tasks raise ValueError, and framework shape or restoration errors propagate.

See Evaluate for held-out scoring, Backends for native inference behavior, or Deploy to archive the selected trained model.