Skip to content

Testing Parameters

Parameters for model evaluation and prediction generation. Used with the .test() and .predict() methods on a scenario model.

For a complete, version-matched MultilabelClassificationTask training and test script that wires both TrainingParams and TestingParams.local_save_location, see Predict Weekly Category Purchases.

test() vs predict()

Both methods accept TestingParams, but they differ in behavior:

test() predict()
Purpose Evaluate against ground truth Score new/unseen data
Metrics Computed and logged Ignored
Data split Uses the test split from training Uses prediction_date as cutoff
prediction_date Optional (defaults to test split boundary) Required for every split type

Use test() during development to measure model quality. Use predict() to generate production predictions.

Sparse targets change the native test denominator

A None target is excluded from native test metric aggregation; EntityIds does not turn it into a miss. If the business denominator must include every requested entity, run predict() for the complete ID file, left-join outcomes, and count entities with no outcome as misses. Report that all-entity metric separately from native test() metrics.

from monad.config import TestingParams, OutputType

TestingParams

Parameter Type Default Description
output_type OutputType required Format in which to save the predictions. See OutputType below.
local_save_location Path \| None None Local file path for predictions in TSV format. Must end with .tsv.
remote_save_location DataLocation \| None None Remote database table for storing predictions. Snowflake and Databricks are supported.
limit_test_batches int \| None None Limit number of test/predict batches to process.
precision Literal[...] 32 Float precision used for testing. See Precision Values.
prediction_date datetime \| None None Date for which to make predictions. Required by predict(), which raises an error when it is missing. In test(), it replaces the test split boundary for entity-based splits and is ignored, with a warning, for time-based splits.

Inherited Parameters

Shared with TrainingParams:

Parameter Type Default Description
show_entity_progress bool False Adds a progress bar that tracks processed entities next to the batch progress bars. Unless limit_test_batches is set, a distinct-entity count runs on global rank 0 before the stage starts to give the bar its total. With the default False, no counting queries run. Has no effect when the ENABLE_PROGRESS_BAR environment variable is False.
devices list[int] \| int \| "auto" "auto" GPU devices to use. "auto" selects the least-occupied GPU automatically. It does not switch to CPU when no GPU is available; set accelerator="cpu" for CPU inference. A positive int specifies count; a list[int] specifies device indices; -1 uses all available GPUs.
accelerator "cpu" \| "gpu" "gpu" Accelerator type.
strategy str \| None None Distributed strategy.
nccl_timeout timedelta \| None None Timeout for NCCL collective operations in distributed inference (DDP/FSDP). A bare number is interpreted as seconds. Ignored on a single device.
rank_sync_timeout timedelta \| None None Timeout for a dedicated per-step rank-synchronization barrier that absorbs data-loading skew between ranks, independently of nccl_timeout (which then only needs to cover the gradient sync). A bare number is interpreted as seconds. Ignored on a single device.
metrics list[MetricParams \| CustomMetric] [] Metrics to compute during testing. See MetricParams and CustomMetric.
top_k int \| None None Limit predictions to top-k items/classes (recommendation, multilabel).
predictions_threshold float \| None None Classification threshold (binary, multilabel). Mutually exclusive with top_k.
entity_ids EntityIds \| None None Limit the run to specific entity IDs. See EntityIds.
callbacks list[Callback] [] PyTorch Lightning callbacks.
approximate_decoding_params ApproximateDecodingParams \| None None Approximate decoding for recommendation tasks.

OutputType

from monad.config import OutputType

# Available values:
OutputType.RAW_MODEL
OutputType.ENCODED
OutputType.DECODED
OutputType.SEMANTIC

Meaning Per Task Type

Task RAW_MODEL ENCODED DECODED SEMANTIC
Binary Logits Logits Probabilities 0 or 1 (based on threshold)
Multiclass Log-softmax Log-softmax Probabilities (with filtering) Class names (with filtering)
Multilabel Logits Logits Probabilities (with filtering) Class names (with filtering, requires top_k)
Regression Raw output Internal representation Human-readable values Human-readable values
Recommendation Raw output Sketch (compact) Probabilities per item Item IDs/names

For all task types except recommendations, we suggest using DECODED.

Tip

Use DECODED for most inference pipelines. Use SEMANTIC when results need to be human-readable. Use ENCODED for recommendation models when you need compact output that can be decoded later with readout_sketch().

EntityIds

Restricts a run to a specific set of entities — for example, to score only active customers, hold out a test population, or debug on a small subset. Passed as entity_ids on TestingParams (and, for training, on TrainingParams).

Field Type Default Description
subquery str \| None None SQL query returning the entity IDs to use.
file str \| None None Path to a text/CSV file whose required header is column0, followed by one entity ID per line.
matching bool True true keeps only the listed IDs (include); false excludes them.

Provide exactly one of subquery or file — supplying both, or neither, raises a validation error. An EntityIds.file is read with read_csv_auto() and selected as column0, so a headerless one-ID-per-line file fails. Use this exact shape:

entity_ids.txt
column0
entity_001
entity_002
from monad.config import EntityIds, OutputType, TestingParams

testing_params = TestingParams(
    output_type=OutputType.DECODED,
    entity_ids=EntityIds(
        subquery='SELECT DISTINCT("customer_id") FROM CUSTOMERS WHERE "age" > 18',
        matching=True,
    ),
)

Usage Examples

Basic Test and Predict

Both methods accept an optional seed parameter for reproducible ordering of results.

from datetime import datetime, timezone
from pathlib import Path
from monad.config import TestingParams, OutputType, MetricParams

testing_params = TestingParams(
    output_type=OutputType.DECODED,
    devices=[0],
    local_save_location=Path("./predictions.tsv"),
    metrics=[
        MetricParams(alias="auroc", metric_name="AUROC", kwargs={"task": "binary"}),
    ],
)

# Test — loads checkpoint, returns test metrics (predictions are saved to the configured location)
results = module.test(testing_params)

# Predict — loads checkpoint, saves predictions to local/remote location
predict_params = TestingParams(
    output_type=OutputType.DECODED,
    devices=[0],
    prediction_date=datetime(2024, 6, 1, tzinfo=timezone.utc),
    local_save_location=Path("./predictions.tsv"),
)
module.predict(predict_params, seed=42)

Recommendation with Top-K

testing_params = TestingParams(
    output_type=OutputType.SEMANTIC,
    devices=[0],
    top_k=10,
    prediction_date=datetime(2024, 6, 1, tzinfo=timezone.utc),
    local_save_location=Path("./top10_predictions.tsv"),
)

module.predict(testing_params)

Multi-GPU Inference

testing_params = TestingParams(
    output_type=OutputType.DECODED,
    devices=[0, 1, 2, 3],
    strategy="ddp",
    prediction_date=datetime(2024, 6, 1, tzinfo=timezone.utc),
    local_save_location=Path("./predictions.tsv"),
)

module.predict(testing_params)

Prediction Output Schema

predict() writes a TSV with header entity_id\tprediction — one row per scored entity; for binary tasks prediction is the positive-class probability. test() additionally includes ground_truth.

Scored population: predict() scores the entities that the data config yields at prediction_date, narrowed by entity_ids when it is set. The target function does not run during prediction, so guards inside it (such as history["transactions"].count() >= 2) do not filter prediction rows. To restrict who is scored, pass EntityIds or filter the data source with where_condition. Entities excluded this way are absent from the file — a row count smaller than the raw entity count is expected, not data loss.

test() also omits an entity when its target function returns None; that is an exclusion, not an empty recommendation target or a miss. For a fixed all-eligible population where no future purchase is a valid outcome, run predict() for those IDs, derive the fixed-window truth independently, left-join it to the predictions, and count every empty future outcome as a miss.

Writing Predictions to Snowflake

from datetime import datetime, timezone
from monad.config import TestingParams, OutputType
from monad.config.data_source import DataLocation

testing_params = TestingParams(
    output_type=OutputType.DECODED,
    devices=[0],
    prediction_date=datetime(2024, 6, 1, tzinfo=timezone.utc),
    remote_save_location=DataLocation(
        database_type="snowflake",
        connection_params={
            "user": "${SNOWFLAKE_USER}",
            "password": "${SNOWFLAKE_PASSWORD}",
            "account": "${SNOWFLAKE_ACCOUNT}",
            "warehouse": "${SNOWFLAKE_WAREHOUSE}",
            "database": "MY_DATABASE",
            "schema": "PUBLIC",
        },
        table_name="predictions_output",
    ),
)

module.predict(testing_params)

Writing Predictions to Databricks

from datetime import datetime, timezone
from monad.config import TestingParams, OutputType
from monad.config.data_source import DataLocation

testing_params = TestingParams(
    output_type=OutputType.DECODED,
    devices=[0],
    prediction_date=datetime(2024, 6, 1, tzinfo=timezone.utc),
    remote_save_location=DataLocation(
        database_type="databricks",
        connection_params={
            "host": "${DATABRICKS_HOST}",
            "warehouse_id": "${DATABRICKS_WAREHOUSE_ID}",
            "token": "${DATABRICKS_TOKEN}",
        },
        table_name="predictions_output",
    ),
)

module.predict(testing_params)

Databricks write behavior

The target table is created on demand if it does not exist, and rows are appended in batches (no surrounding transaction). Tune the batch size with the DATABRICKS_WRITE_BATCH_SIZE environment variable (default 1000).


TSV Output Schema

When local_save_location is set, output is a tab-separated file. predict() writes entity_id and prediction; test() adds ground_truth. The native evaluate_predictions() function consumes the prediction and ground_truth columns rather than task-specific score_* or label_* columns.

The values inside those two columns depend on the task and output type:

Task prediction ground_truth from test()
Binary scalar score or decoded probability scalar 0 or 1
Multiclass comma-separated score vector over all classes matching comma-separated class distribution
Multilabel comma-separated score vector over all classes matching comma-separated 0/1 flags
Regression scalar value matching scalar value
Recommendation comma-separated ranked items comma-separated relevant items

Use OutputType.DECODED for the numeric classification/regression forms accepted by the native evaluator. Recommendation evaluation consumes the ranked item lists. predict() output has no ground_truth and therefore cannot be passed directly to evaluate_predictions().

For a regression task with multiple targets, select one target and write scalar prediction and ground_truth columns before calling the native evaluator; it evaluates one continuous target at a time.


Prediction Utilities

readout_sketch()

Decode recommendation predictions saved with OutputType.ENCODED into per-item scores.

from monad.ui.module import readout_sketch

generator = readout_sketch(
    predictions_file="./predictions.tsv",
    checkpoint_path="./reco_model",
)

for entity_id, scores in generator:
    print(f"Entity: {entity_id}, Scores shape: {scores.shape}")
Parameter Type Description
predictions_file str Path to predictions file saved with OutputType.ENCODED.
checkpoint_path str Path to the recommendation model checkpoint.

Returns a generator yielding (entity_id: str, scores: np.ndarray) tuples.

read_target_entity_ids()

Get the mapping from target entity IDs (e.g., product IDs) to their indices in the decoded sketch.

from monad.ui.module import read_target_entity_ids

target_to_index = read_target_entity_ids(
    checkpoint_path="./reco_model",
)
# Returns: {"product_001": 0, "product_002": 1, ...}
Parameter Type Description
checkpoint_path str Path to the recommendation model checkpoint.

Returns a dict[str, int] mapping entity IDs to indices.