setu.validate_estimator#

setu.validate_estimator(estimator, val_data, key, *, training_result=None, checks=None, strict=False, stratify_by=None, config=None)#
Overloads:
  • estimator (NLE | MixedNLE | DiscreteNLE | NRE), val_data (ValidationDataset), key (Array), stratify_by (Array), training_result (AnyTrainingResult | None), checks (list[str] | None), strict (bool), config (ValidationConfig | None) → dict[int, ValidationResult]

  • estimator (NLE | MixedNLE | DiscreteNLE | NRE), val_data (ValidationDataset), key (Array), stratify_by (None), training_result (AnyTrainingResult | None), checks (list[str] | None), strict (bool), config (ValidationConfig | None) → ValidationResult

Parameters:
Return type:

ValidationResult | dict[int, ValidationResult]

Validate a trained estimator (NLE, MixedNLE, DiscreteNLE, NRE) before inference.

Runs comprehensive validation checks. For NLE/MixedNLE: training diagnostics, C2ST, MMD, Sliced-Wasserstein, Likelihood SBC, and Likelihood Accuracy. For DiscreteNLE: training diagnostics, C2ST, Likelihood SBC, Likelihood Accuracy, and Total Variation (skips MMD/SW, which assume continuous data). For NRE: training diagnostics, Likelihood Accuracy, Normalizing Constant, and Ratio SBC. The appropriate checks are selected automatically based on the estimator type.

Parameters:
  • estimator (NLE | MixedNLE | DiscreteNLE | NRE) – Trained NLE, MixedNLE, DiscreteNLE, or NRE to validate

  • val_data (ValidationDataset) – Validation dataset (MUST be separate from training data)

  • key (Array) – JAX random key

  • training_result (TrainingResult | MixedNLETrainingResult | NRETrainingResult | DiscreteNLETrainingResult | None) – Optional training result for diagnostics

  • checks (list[str] | None) – List of checks to run. Defaults depend on estimator type: NLE/MixedNLE: [“training”, “c2st”, “mmd”, “sw”, “sbc”, “la”]; DiscreteNLE: [“training”, “c2st”, “sbc”, “la”, “tv”]; NRE: [“training”, “la”, “nc”, “ratio_sbc”]. Set to a subset to run specific checks only.

  • strict (bool) – If True, raise ValidationError on any failure

  • stratify_by (Array | None) – Optional integer labels, shape (n,). If provided, validation is run independently for each unique label and a dict[int, ValidationResult] is returned instead.

  • config (ValidationConfig | None) – Thresholds and sample counts for each check. Defaults to ValidationConfig() with sensible defaults.

Returns:

ValidationResult with all diagnostics, or dict[int, ValidationResult] when stratify_by is provided.

Raises:
Return type:

ValidationResult | dict[int, ValidationResult]

Example

>>> val_data = ValidationDataset(theta_val, x_val)
>>> result = validate_nle(nle, val_data, key, training_result=train_result)
>>> # Custom thresholds:
>>> cfg = ValidationConfig(c2st_fail=0.50, mmd_warn=0.05)
>>> result = validate_nle(nle, val_data, key, config=cfg)