setu.validate_nle#
- setu.validate_nle(estimator, val_data, key, *, training_result=None, checks=None, strict=False, stratify_by=None, config=None)[source]#
- 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:
estimator (NLE | MixedNLE | DiscreteNLE | NRE)
val_data (ValidationDataset)
key (Array)
training_result (TrainingResult | MixedNLETrainingResult | NRETrainingResult | DiscreteNLETrainingResult | None)
strict (bool)
stratify_by (Array | None)
config (ValidationConfig | None)
- Return type:
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 validateval_data (
ValidationDataset) – Validation dataset (MUST be separate from training data)key (
Array) – JAX random keytraining_result (
TrainingResult|MixedNLETrainingResult|NRETrainingResult|DiscreteNLETrainingResult|None) – Optional training result for diagnosticschecks (
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 failurestratify_by (
Array|None) – Optional integer labels, shape(n,). If provided, validation is run independently for each unique label and adict[int, ValidationResult]is returned instead.config (
ValidationConfig|None) – Thresholds and sample counts for each check. Defaults toValidationConfig()with sensible defaults.
- Returns:
ValidationResult with all diagnostics, or
dict[int, ValidationResult]whenstratify_byis provided.- Raises:
ValidationError – If strict=True and any check fails
ValueError – If val_data dimensions don’t match estimator
- Return type:
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)