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.
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:
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.