From 958ff10c31f320c820e6abfaef6b11b4e783bf6a Mon Sep 17 00:00:00 2001 From: Jennifer Pollack Date: Wed, 30 Sep 2026 18:32:27 +0200 Subject: [PATCH 1/3] Refactor quality metrics to expose multiple diagnostics --- src/wf_psf/quality_control/metrics/base.py | 8 ++++++-- src/wf_psf/quality_control/metrics/goodness_of_fit.py | 4 ++++ src/wf_psf/quality_control/metrics/pixel_masks.py | 9 +++++++++ 3 files changed, 19 insertions(+), 2 deletions(-) diff --git a/src/wf_psf/quality_control/metrics/base.py b/src/wf_psf/quality_control/metrics/base.py index d3bc0783..38f73e9a 100644 --- a/src/wf_psf/quality_control/metrics/base.py +++ b/src/wf_psf/quality_control/metrics/base.py @@ -13,7 +13,7 @@ from abc import ABC, abstractmethod import numpy as np -from typing import Any +from typing import Any, ClassVar from wf_psf.quality_control.context import QualityControlContext @@ -27,6 +27,9 @@ class QualityMetric(ABC): Unique identifier for the metric implementation. Used by the MetricsRegistry to register and retrieve metric classes. + diagnostics : ClassVar[frozenset[str]] + Immutable set of diagnostic names exposed by the metric. + params : dict[str, Any] Parameter set for configuring a specific metric. @@ -37,7 +40,8 @@ class QualityMetric(ABC): """ - name: str + name: ClassVar[str] + diagnostics: ClassVar[frozenset[str]] def __init__(self, params: dict[str, Any]): self.params = params diff --git a/src/wf_psf/quality_control/metrics/goodness_of_fit.py b/src/wf_psf/quality_control/metrics/goodness_of_fit.py index 2d8df453..958880dc 100644 --- a/src/wf_psf/quality_control/metrics/goodness_of_fit.py +++ b/src/wf_psf/quality_control/metrics/goodness_of_fit.py @@ -12,12 +12,16 @@ from .base import QualityMetric import numpy as np +from typing import ClassVar class GoodnessOfFitMetric(QualityMetric): """Compute a goodness-of-fit metric (e.g. reduced chi square) for each dataset sample.""" name = "goodness_of_fit" + diagnostics: ClassVar[frozenset[str]] = frozenset( + {"chi_square", "reduced_chi_square"} + ) def compute(self, context: QualityControlContext) -> dict[str, np.ndarray]: """Compute reduced chi-square values for each dataset sample.""" diff --git a/src/wf_psf/quality_control/metrics/pixel_masks.py b/src/wf_psf/quality_control/metrics/pixel_masks.py index 10d67793..841ff4c7 100644 --- a/src/wf_psf/quality_control/metrics/pixel_masks.py +++ b/src/wf_psf/quality_control/metrics/pixel_masks.py @@ -8,12 +8,21 @@ """ from .base import QualityMetric +from typing import ClassVar class PixelMaskMetric(QualityMetric): """Evaluate pixel-mask metrics for each dataset sample.""" name = "pixel_mask" + diagnostics: ClassVar[frozenset[str]] = frozenset( + { + "total_masked_pixels", + "total_masked_fraction", + "aperture_masked_pixels", + "aperture_masked_fraction", + } + ) def compute(self, dataset): """Compute pixel-mask metrics for each dataset sample.""" From 31eb517779c4e2132e7e44a0c324ab13334b949d Mon Sep 17 00:00:00 2001 From: Jennifer Pollack Date: Wed, 30 Sep 2026 18:45:22 +0200 Subject: [PATCH 2/3] Update configuration for diagnostic selection - Add diagnostic attribute to RejectionPolicyConfig - Update rejection config section parser - Remove validation methods added to pipeline.py - Update config_test.py with new and improved test cases - Update fixtures and remove deprecated YAML fixtures --- config/quality_control.yaml | 4 +- src/wf_psf/quality_control/config.py | 100 ++------ .../tests/test_quality_control/config_test.py | 232 ++++++++++-------- .../tests/test_quality_control/conftest.py | 1 + .../data/invalid/invalid_metric.yaml | 3 - .../data/invalid/metric_missing_sections.yaml | 2 - .../metric_resource_identifier_unknown.yaml | 15 -- .../invalid/rejection_metric_not_enabled.yaml | 26 -- .../data/valid/quality_control.yaml | 4 +- ...y_control_multiple_rejection_policies.yaml | 3 + ...ity_control_rejection_policy_disabled.yaml | 52 ++++ 11 files changed, 207 insertions(+), 235 deletions(-) delete mode 100644 src/wf_psf/tests/test_quality_control/data/invalid/invalid_metric.yaml delete mode 100644 src/wf_psf/tests/test_quality_control/data/invalid/metric_missing_sections.yaml delete mode 100644 src/wf_psf/tests/test_quality_control/data/invalid/metric_resource_identifier_unknown.yaml delete mode 100644 src/wf_psf/tests/test_quality_control/data/invalid/rejection_metric_not_enabled.yaml create mode 100644 src/wf_psf/tests/test_quality_control/data/valid/quality_control_rejection_policy_disabled.yaml diff --git a/config/quality_control.yaml b/config/quality_control.yaml index 8517e663..86341e5c 100644 --- a/config/quality_control.yaml +++ b/config/quality_control.yaml @@ -13,6 +13,7 @@ metrics: params: aperture: type: circular + centre: stamp_centre radius: 2.6 unit: sigma @@ -21,7 +22,6 @@ metrics: required_resources: - psf_models.standard params: - statistic: reduced_chi_square normalize_residuals: true @@ -29,12 +29,14 @@ rejection: pixel_mask: enabled: true + diagnostic: aperture_masked_fraction policy: threshold: value: 0.25 goodness_of_fit: enabled: false + diagnostic: reduced_chi_square policy: threshold: value: 3.0 diff --git a/src/wf_psf/quality_control/config.py b/src/wf_psf/quality_control/config.py index 32943780..a6c6d76b 100644 --- a/src/wf_psf/quality_control/config.py +++ b/src/wf_psf/quality_control/config.py @@ -48,11 +48,15 @@ class RejectionPolicyConfig: enabled : bool Whether rejection policy is enabled. + diagnostic : str | None + Diagnostic from the quality metric to use when applying a rejection policy. + policy : dict[str, Any] Rejection policy configuration keyed by policy type. """ enabled: bool = False + diagnostic: str | None = None policy: dict[str, Any] = field(default_factory=dict) @@ -236,6 +240,14 @@ def parse_rejection_policy_config( policies[metric_name] = RejectionPolicyConfig(enabled=False) continue + diagnostic = cfg.get("diagnostic", None) + + if not isinstance(diagnostic, str) or not diagnostic: + raise ValueError( + f"Rejection policy configuration for '{metric_name}' " + "must specify a non-empty `diagnostic`." + ) + if "policy" not in cfg: raise ValueError( f"Rejection policy configuration for '{metric_name}' " @@ -268,7 +280,7 @@ def parse_rejection_policy_config( ) policies[metric_name] = RejectionPolicyConfig( - enabled=enabled, policy=dict(policy) + enabled=enabled, diagnostic=diagnostic, policy=dict(policy) ) return policies @@ -332,80 +344,6 @@ def parse_resources_config( return ResourcesConfig(available=dict(config)) -# validators for internal consistency of config sections -def validate_quality_control_config(config: QualityControlConfig) -> None: - """Validate internal consistency of a quality control configuration. - - Parameters - ---------- - config : QualityControlConfig - Parsed quality control configuration. - - Raises - ------ - ValueError - If any cross-section configuration dependency is invalid. - """ - validate_metric_resources(config) - validate_rejection_policy_metrics(config) - - -def validate_metric_resources(config: QualityControlConfig) -> None: - """Validate that all metric resource requirements can be resolved. - - Parameters - ---------- - config : QualityControlConfig - Parsed quality control configuration. - - Raises - ------ - ValueError - If a required resource identifier is not available in the configured resources. - """ - for metric_name, metric in config.metrics.items(): - for resource_id in metric.required_resources: - resources = config.resources.available - - if ( - resource_id.family not in resources - or resource_id.variant not in resources[resource_id.family] - ): - raise ValueError( - f"Metric '{metric_name}' requires unknown resource '{resource_id}'." - ) - - -def validate_rejection_policy_metrics(config: QualityControlConfig) -> None: - """Validate rejection policies against configured quality metrics. - - Parameters - ---------- - config : QualityControlConfig - Parsed quality control configuration. - - Raises - ------ - ValueError - If an enabled rejection policy references an unknown or disabled - quality metric. - """ - for metric_name, rejection_policy in config.rejection.items(): - if not rejection_policy.enabled: - continue - - if metric_name not in config.metrics: - raise ValueError( - f"Rejection policy configured for unknown metric '{metric_name}'." - ) - - if not config.metrics[metric_name].enabled: - raise ValueError( - f"Rejection policy cannot be enabled because metric " - f"'{metric_name}' is disabled." - ) - - SECTION_PARSERS = { "metrics": parse_metrics_config, "rejection": parse_rejection_policy_config, @@ -437,7 +375,7 @@ def load(self) -> QualityControlConfig: Returns ------- QualityControlConfig - Parsed and validated quality control configuration. + Parsed quality control configuration. Raises ------ @@ -445,8 +383,8 @@ def load(self) -> QualityControlConfig: If a configuration section has an invalid structure or type. ValueError - If the parsed configuration contains inconsistent references - between metrics, resources, or rejection policies. + If a configuration section contains an invalid value or + configuration structure. """ qc_config = read_yaml(self.qc_config_path) config = {} @@ -455,8 +393,4 @@ def load(self) -> QualityControlConfig: values = qc_config.get(section, {}) config[section] = parser(values) - qc = QualityControlConfig(**config) - - validate_quality_control_config(qc) - - return qc + return QualityControlConfig(**config) diff --git a/src/wf_psf/tests/test_quality_control/config_test.py b/src/wf_psf/tests/test_quality_control/config_test.py index 8789b67b..09af1b9a 100644 --- a/src/wf_psf/tests/test_quality_control/config_test.py +++ b/src/wf_psf/tests/test_quality_control/config_test.py @@ -6,7 +6,6 @@ """ -from contextlib import nullcontext as does_not_raise from pathlib import Path import pytest from wf_psf.quality_control.config import ( @@ -20,8 +19,6 @@ from wf_psf.quality_control.config import ( parse_resources_config, parse_rejection_policy_config, - validate_metric_resources, - validate_rejection_policy_metrics, ) from wf_psf.quality_control.resource_identifier import ResourceIdentifier @@ -35,6 +32,16 @@ def load_config(config_file: str) -> QualityControlConfig: def test_quality_control_config_loading(): config = load_config("valid/quality_control.yaml") + pixel_mask_params = { + "aperture": { + "type": "circular", + "centre": "stamp_centre", + "radius": 2.6, + "unit": "sigma", + }, + } + + assert isinstance(config, QualityControlConfig) assert isinstance(config.resources, ResourcesConfig) assert "standard" in config.resources.available["psf_models"] assert "oversampled" in config.resources.available["psf_models"] @@ -50,12 +57,16 @@ def test_quality_control_config_loading(): assert "pixel_mask" in config.metrics assert isinstance(config.metrics["pixel_mask"], QualityMetricConfig) assert config.metrics["pixel_mask"].enabled is True + assert config.metrics["pixel_mask"].params == pixel_mask_params assert config.metrics["pixel_mask"].required_resources == [] assert isinstance(config.metrics["goodness_of_fit"], QualityMetricConfig) assert config.metrics["goodness_of_fit"].required_resources == ( [ResourceIdentifier.from_string("psf_models.standard")] ) + assert config.metrics["goodness_of_fit"].params == { + "normalize_residuals": True, + } assert isinstance(config.rejection["pixel_mask"], RejectionPolicyConfig) assert config.rejection["pixel_mask"].policy == { @@ -122,7 +133,33 @@ def test_metrics_configuration_must_be_mapping(): ## Tests for parsing rejection policy configurations -def test_rejection_configuration_must_be_mapping(): +def test_rejection_policy_configuration_valid(): + result = parse_rejection_policy_config( + { + "goodness_of_fit": { + "enabled": True, + "diagnostic": "reduced_chi_square", + "policy": { + "threshold": { + "value": 0.25, + } + }, + } + } + ) + + assert result["goodness_of_fit"] == RejectionPolicyConfig( + enabled=True, + diagnostic="reduced_chi_square", + policy={ + "threshold": { + "value": 0.25, + } + }, + ) + + +def test_rejection_policy_configuration_must_be_mapping(): with pytest.raises( TypeError, match="Rejection policy configuration must be a mapping" ): @@ -137,20 +174,77 @@ def test_rejection_policy_configuration_metric_must_be_mapping(): parse_rejection_policy_config({"goodness_of_fit": 0.25}) -def test_rejection_policy_must_specify_policy(): +def test_rejection_policy_enabled_defaults_to_false(): + policies = parse_rejection_policy_config( + { + "goodness_of_fit": { + "diagnostic": "reduced_chi_square", + "policy": {"threshold": {"value": 0.25}}, + } + } + ) + + assert policies["goodness_of_fit"] == RejectionPolicyConfig(enabled=False) + + +def test_rejection_policy_metric_enabled_flag_must_be_boolean(): with pytest.raises( - ValueError, - match="must specify a `policy`", + TypeError, + match="Rejection policy `enabled` flag for 'goodness_of_fit' must be boolean.", ): parse_rejection_policy_config( { "goodness_of_fit": { - "enabled": True, + "enabled": "foo", + "diagnostic": "reduced_chi_square", + "policy": "threshold", } } ) +def test_rejection_policy_disabled_policies_are_skipped(): + rejection_policy = { + "goodness_of_fit": { + "enabled": False, + "diagnostic": "None", + "policy": "not a mapping", + }, + } + + policies = parse_rejection_policy_config(rejection_policy) + + assert policies["goodness_of_fit"].enabled is False + assert policies["goodness_of_fit"].diagnostic is None + assert policies["goodness_of_fit"].policy == {} + + +def test_rejection_policy_must_specify_non_empty_diagnostic(): + rejection_policy = { + "goodness_of_fit": { + "enabled": True, + "diagnostic": None, + "policy": "not a mapping", + }, + } + + with pytest.raises( + ValueError, + match="Rejection policy configuration for 'goodness_of_fit' must specify a non-empty `diagnostic`.", + ): + parse_rejection_policy_config(rejection_policy) + + +def test_rejection_policy_must_specify_policy(): + with pytest.raises( + ValueError, + match="must specify a `policy`", + ): + parse_rejection_policy_config( + {"goodness_of_fit": {"enabled": True, "diagnostic": "reduced_chi_square"}} + ) + + def test_rejection_policy_must_be_mapping(): with pytest.raises( TypeError, @@ -160,6 +254,7 @@ def test_rejection_policy_must_be_mapping(): { "goodness_of_fit": { "enabled": True, + "diagnostic": "reduced_chi_square", "policy": "threshold", } } @@ -185,120 +280,49 @@ def test_rejection_policy_must_specify_exactly_one_policy(policy): { "goodness_of_fit": { "enabled": True, + "diagnostic": "reduced_chi_square", "policy": policy, } } ) -def test_disabled_rejection_policy_does_not_require_policy(): - result = parse_rejection_policy_config({"goodness_of_fit": {"enabled": False}}) - - assert result == {"goodness_of_fit": RejectionPolicyConfig(enabled=False)} - - -def test_rejection_policy_configuration(): - result = parse_rejection_policy_config( - { - "goodness_of_fit": { - "enabled": True, - "policy": { - "threshold": { - "value": 0.25, - } - }, - } +def test_rejection_policy_identifier_must_be_string(): + rejection_policy = { + "goodness_of_fit": { + "enabled": True, + "diagnostic": "reduced_chi_square", + "policy": {123: {}}, } - ) - - assert result["goodness_of_fit"] == RejectionPolicyConfig( - enabled=True, - policy={ - "threshold": { - "value": 0.25, - } - }, - ) - + } -## Tests for parsing reporting configurations -def test_reporting_configuration_must_be_mapping(): with pytest.raises( TypeError, - match="Reporting configuration must be a mapping", + match="Rejection policy identifier '123' for 'goodness_of_fit' must be a string.", ): - load_config("invalid/reporting_invalid_type.yaml") - - -# Tests for validation methods -def test_validate_metric_resources_all_valid(qc_config_factory): - with does_not_raise(): - validate_metric_resources(qc_config_factory()) + parse_rejection_policy_config(rejection_policy) -@pytest.mark.parametrize( - "required_resource", - [ - ResourceIdentifier.from_string("images.segmentation_maps"), - ResourceIdentifier.from_string("psf_models.imaginary"), - ], -) -def test_validate_metric_resources_unknown_resource( - qc_config_factory, - required_resource, -): - config = qc_config_factory(required_resources=[required_resource]) - - with pytest.raises( - ValueError, - match=( - f"Metric 'goodness_of_fit' requires unknown resource '{required_resource}'." - ), - ): - validate_metric_resources(config) - - -def test_validate_rejection_policy_metrics_all_valid(qc_config_factory): - with does_not_raise(): - validate_rejection_policy_metrics(qc_config_factory()) - - -def test_validate_rejection_policy_metrics_metric_not_found(qc_config_factory): - config = qc_config_factory(rejection_metric="pixel_mask") - - with pytest.raises( - ValueError, - match="Rejection policy configured for unknown metric 'pixel_mask'.", - ): - validate_rejection_policy_metrics(config) - - -def test_validate_rejection_policy_metrics_metric_not_enabled(qc_config_factory): - metric = { - "goodness_of_fit": QualityMetricConfig( - enabled=False, - required_resources=[], - ) +def test_rejection_policy_params_must_be_mapping(): + rejection_policy = { + "goodness_of_fit": { + "enabled": True, + "diagnostic": "reduced_chi_square", + "policy": {"threshold": 3}, + } } - config = qc_config_factory(metrics=metric) - with pytest.raises( - ValueError, - match="Rejection policy cannot be enabled because metric 'goodness_of_fit' is disabled.", + TypeError, + match="Rejection policy parameters for 'goodness_of_fit' must be a mapping.", ): - validate_rejection_policy_metrics(config) - + parse_rejection_policy_config(rejection_policy) -# Integration tests -def test_load_config_validates_configuration_pass(): - with does_not_raise(): - load_config("valid/quality_control.yaml") - -def test_load_config_validates_configuration_raise_unknown_identifier(): +## Tests for parsing reporting configurations +def test_reporting_configuration_must_be_mapping(): with pytest.raises( - ValueError, - match="Metric 'goodness_of_fit' requires unknown resource", + TypeError, + match="Reporting configuration must be a mapping", ): - load_config("invalid/metric_resource_identifier_unknown.yaml") + load_config("invalid/reporting_invalid_type.yaml") diff --git a/src/wf_psf/tests/test_quality_control/conftest.py b/src/wf_psf/tests/test_quality_control/conftest.py index 7811f3d4..fdf9ad19 100644 --- a/src/wf_psf/tests/test_quality_control/conftest.py +++ b/src/wf_psf/tests/test_quality_control/conftest.py @@ -38,6 +38,7 @@ def factory( rejection_default = { rejection_metric or "goodness_of_fit": RejectionPolicyConfig( enabled=True, + diagnostic="reduced_chi_square", policy={ "threshold": { "value": 0.25, diff --git a/src/wf_psf/tests/test_quality_control/data/invalid/invalid_metric.yaml b/src/wf_psf/tests/test_quality_control/data/invalid/invalid_metric.yaml deleted file mode 100644 index 35ee0947..00000000 --- a/src/wf_psf/tests/test_quality_control/data/invalid/invalid_metric.yaml +++ /dev/null @@ -1,3 +0,0 @@ -metrics: - - pixel_mask: true diff --git a/src/wf_psf/tests/test_quality_control/data/invalid/metric_missing_sections.yaml b/src/wf_psf/tests/test_quality_control/data/invalid/metric_missing_sections.yaml deleted file mode 100644 index 1cb23f13..00000000 --- a/src/wf_psf/tests/test_quality_control/data/invalid/metric_missing_sections.yaml +++ /dev/null @@ -1,2 +0,0 @@ -metrics: - pixel_mask: \ No newline at end of file diff --git a/src/wf_psf/tests/test_quality_control/data/invalid/metric_resource_identifier_unknown.yaml b/src/wf_psf/tests/test_quality_control/data/invalid/metric_resource_identifier_unknown.yaml deleted file mode 100644 index 50023177..00000000 --- a/src/wf_psf/tests/test_quality_control/data/invalid/metric_resource_identifier_unknown.yaml +++ /dev/null @@ -1,15 +0,0 @@ -resources: - - psf_models: - standard: - inference_config: inference_standard.yaml - oversampled: - inference_config: inference_oversampled.yaml - -metrics: - - goodness_of_fit: - enabled: true - required_resources: - - images.segmentation_maps - diff --git a/src/wf_psf/tests/test_quality_control/data/invalid/rejection_metric_not_enabled.yaml b/src/wf_psf/tests/test_quality_control/data/invalid/rejection_metric_not_enabled.yaml deleted file mode 100644 index bdf0b460..00000000 --- a/src/wf_psf/tests/test_quality_control/data/invalid/rejection_metric_not_enabled.yaml +++ /dev/null @@ -1,26 +0,0 @@ -resources: - - psf_models: - standard: - inference_config: inference_standard.yaml - oversampled: - inference_config: inference_oversampled.yaml - -metrics: - - pixel_mask: - enabled: false - params: - aperture: - type: circular - radius: 2.6 - unit: sigma - - -rejection: - - pixel_mask: - enabled: true - policy: - threshold: - value: 0.25 diff --git a/src/wf_psf/tests/test_quality_control/data/valid/quality_control.yaml b/src/wf_psf/tests/test_quality_control/data/valid/quality_control.yaml index da49a6d3..fd03eb5d 100644 --- a/src/wf_psf/tests/test_quality_control/data/valid/quality_control.yaml +++ b/src/wf_psf/tests/test_quality_control/data/valid/quality_control.yaml @@ -13,6 +13,7 @@ metrics: params: aperture: type: circular + centre: stamp_centre radius: 2.6 unit: sigma @@ -21,7 +22,6 @@ metrics: required_resources: - psf_models.standard params: - statistic: reduced_chi_square normalize_residuals: true shapes: @@ -34,12 +34,14 @@ rejection: pixel_mask: enabled: true + diagnostic: aperture_masked_fraction policy: threshold: value: 3.0 goodness_of_fit: enabled: false + diagnostic: reduced_chi_square policy: threshold: value: 3.0 diff --git a/src/wf_psf/tests/test_quality_control/data/valid/quality_control_multiple_rejection_policies.yaml b/src/wf_psf/tests/test_quality_control/data/valid/quality_control_multiple_rejection_policies.yaml index 555b0c25..612c1dd6 100644 --- a/src/wf_psf/tests/test_quality_control/data/valid/quality_control_multiple_rejection_policies.yaml +++ b/src/wf_psf/tests/test_quality_control/data/valid/quality_control_multiple_rejection_policies.yaml @@ -13,6 +13,7 @@ metrics: params: aperture: type: circular + centre: stamp_centre radius: 2.6 unit: sigma @@ -27,12 +28,14 @@ metrics: rejection: pixel_mask: + diagnostic: aperture_masked_fraction enabled: true policy: threshold: value: 3.0 goodness_of_fit: + diagnostic: reduced_chi_square enabled: true policy: threshold: diff --git a/src/wf_psf/tests/test_quality_control/data/valid/quality_control_rejection_policy_disabled.yaml b/src/wf_psf/tests/test_quality_control/data/valid/quality_control_rejection_policy_disabled.yaml new file mode 100644 index 00000000..35beb77a --- /dev/null +++ b/src/wf_psf/tests/test_quality_control/data/valid/quality_control_rejection_policy_disabled.yaml @@ -0,0 +1,52 @@ +resources: + + psf_models: + standard: + inference_config: inference_standard.yaml + oversampled: + inference_config: inference_oversampled.yaml + +metrics: + + pixel_mask: + enabled: true + params: + aperture: + type: circular + centre: stamp_centre + radius: 2.6 + unit: sigma + + goodness_of_fit: + enabled: true + required_resources: + - psf_models.standard + params: + normalize_residuals: true + + shapes: + enabled: false + required_resources: + - psf_models.oversampled + + +rejection: + + pixel_mask: + enabled: false + diagnostic: aperture_masked_fraction + policy: + threshold: + value: 3.0 + + goodness_of_fit: + enabled: false + diagnostic: reduced_chi_square + policy: + threshold: + value: 3.0 + +reporting: + + save_metrics: true + log_statistics: true \ No newline at end of file From 89f79a28ec2a2ae2351f45dc9ad73fa3c157b195 Mon Sep 17 00:00:00 2001 From: Jennifer Pollack Date: Wed, 30 Sep 2026 19:00:02 +0200 Subject: [PATCH 3/3] Update pipeline processing - Add validation for metric resource requirements and rejection policy configuration - Call configuration validation from the pipeline constructor - Update pipeline execution to consume the `QualityMetric.compute` diagnostics API - Apply rejection policies using their configured diagnostic - Move cross-section validation tests from config_test.py to pipeline_test.py - Update and extend pipeline test coverage for the revised validation and updated processing flow --- src/wf_psf/quality_control/pipeline.py | 98 ++++- .../test_quality_control/pipeline_test.py | 344 ++++++++++++++++-- 2 files changed, 407 insertions(+), 35 deletions(-) diff --git a/src/wf_psf/quality_control/pipeline.py b/src/wf_psf/quality_control/pipeline.py index 07e7b0d6..75710f5c 100644 --- a/src/wf_psf/quality_control/pipeline.py +++ b/src/wf_psf/quality_control/pipeline.py @@ -34,7 +34,7 @@ class QualityControlResult: Attributes ---------- metrics - Computed quality metrics indexed by metric name. + Computed quality diagnostics indexed by metric name and diagnostic name. validity_masks Boolean validity masks produced by each rejection policy. @@ -44,7 +44,7 @@ class QualityControlResult: rejection policies. """ - metrics: dict[str, np.ndarray] + metrics: dict[str, dict[str, np.ndarray]] validity_masks: dict[str, np.ndarray] @@ -66,6 +66,7 @@ def __init__(self, qc_config_path): self.config = QualityControlConfigHandler(qc_config_path).load() self.metrics_registry = build_metrics_registry() self.rejection_registry = build_rejection_policy_registry() + self.validate_configuration() def _instantiate_metrics(self) -> dict[str, QualityMetric]: """Instantiate enabled quality metric implementations from configuration. @@ -141,6 +142,82 @@ def _resolve_resources(self, provided_resources): resource_manager = Resources(self.config) return resource_manager.resolve(provided_resources) + def validate_configuration(self) -> None: + """Validate internal consistency of a quality control configuration. + + Raises + ------ + ValueError + If any cross-section configuration dependency is invalid. + """ + self.validate_metric_resource_requirements() + self.validate_rejection_policy_metrics() + + def validate_metric_resource_requirements(self) -> None: + """Validate that resource requirements of enabled metrics are configured. + + Raises + ------ + ValueError + If a required resource identifier is not available in the configured resources. + + Notes + ----- + A configured resource can be overriden by the pipeline + caller using the `provided_resources` argument. This static validation ensures that + each resource identifier required by an enabled metric is declared in the + resources configuration, regardless of whethe resource will be ultimately supplied + by the caller or prepared by the pipeline. + + """ + resources = self.config.resources.available + + for metric_name, metric in self.config.metrics.items(): + if not metric.enabled: + continue + + for resource_id in metric.required_resources: + if ( + resource_id.family not in resources + or resource_id.variant not in resources[resource_id.family] + ): + raise ValueError( + f"Metric '{metric_name}' requires unknown resource '{resource_id}'." + ) + + def validate_rejection_policy_metrics(self) -> None: + """Validate rejection policies against configured quality metrics. + + Raises + ------ + ValueError + If an enabled rejection policy references an unknown or disabled + quality metric, or specifies an invalid diagnostic. + """ + for metric_name, rejection_policy in self.config.rejection.items(): + if not rejection_policy.enabled: + continue + + if metric_name not in self.config.metrics: + raise ValueError( + f"Rejection policy configured for unknown metric '{metric_name}'." + ) + + if not self.config.metrics[metric_name].enabled: + raise ValueError( + f"Rejection policy cannot be enabled because metric '{metric_name}' is disabled." + ) + + metric_cls = self.metrics_registry.get(metric_name) + + if metric_cls is None: + raise ValueError(f"Quality metric '{metric_name}' is not registered.") + + if rejection_policy.diagnostic not in metric_cls.diagnostics: + raise ValueError( + f"Diagnostic '{rejection_policy.diagnostic}' is not provided by the metric '{metric_name}'." + ) + def run(self, dataset, provided_resources=None): """Run quality control pipeline. @@ -170,18 +247,23 @@ def run(self, dataset, provided_resources=None): rejection_policies = self._instantiate_rejection_policies() - validity_masks = { - name: policy.apply(metric_results[name]) - for name, policy in rejection_policies.items() - } + validity_masks = {} + for name, policy in rejection_policies.items(): + diagnostic_identifier = self.config.rejection[name].diagnostic + assert diagnostic_identifier is not None + + diagnostic = metric_results[name][diagnostic_identifier] + validity_masks[name] = policy.apply(diagnostic) - if validity_masks: + if validity_masks != {}: # True indicates a valid sample. A sample is valid only if it passes # every enabled rejection policy. valid_mask = np.logical_and.reduce(list(validity_masks.values())) else: + # If rejection policy not enabled, generate boolean unity mask metric_result = next(iter(metric_results.values())) - valid_mask = np.ones(metric_result.shape, dtype=bool) + diagnostic_result = next(iter(metric_result.values())) + valid_mask = np.ones(diagnostic_result.shape, dtype=bool) return QualityControlResult( metrics=metric_results, diff --git a/src/wf_psf/tests/test_quality_control/pipeline_test.py b/src/wf_psf/tests/test_quality_control/pipeline_test.py index 1ff3bb44..9ea2403f 100644 --- a/src/wf_psf/tests/test_quality_control/pipeline_test.py +++ b/src/wf_psf/tests/test_quality_control/pipeline_test.py @@ -1,13 +1,20 @@ +from contextlib import nullcontext as does_not_raise import numpy as np from pathlib import Path import pytest -from unittest.mock import patch +from unittest.mock import call, patch from wf_psf.quality_control.pipeline import QualityControlPipeline -from wf_psf.quality_control.config import QualityControlConfig +from wf_psf.quality_control.config import ( + QualityControlConfig, + QualityMetricConfig, + RejectionPolicyConfig, +) + from wf_psf.quality_control.metrics.pixel_masks import PixelMaskMetric from wf_psf.quality_control.metrics.goodness_of_fit import GoodnessOfFitMetric from wf_psf.quality_control.rejection.threshold import ThresholdRejectionPolicy +from wf_psf.quality_control.resource_identifier import ResourceIdentifier @pytest.fixture @@ -32,6 +39,10 @@ def test_pipeline_constructor(pipeline_factory): # Check rejection registry assert pipeline.rejection_registry.get("threshold") is ThresholdRejectionPolicy + # Check config validation + with does_not_raise(): + pipeline.validate_configuration() + def test_pipeline_instantiate_metrics_valid(pipeline_factory): pipeline = pipeline_factory("valid/quality_control.yaml") @@ -62,21 +73,216 @@ def test_pipeline_instantiate_rejection_policy_valid(pipeline_factory): assert "goodness_of_fit" not in rejection_policies +# Tests for validation methods +def test_validate_metric_resources_all_valid(pipeline_factory): + pipeline = pipeline_factory("valid/quality_control.yaml") + with does_not_raise(): + pipeline.validate_metric_resource_requirements() + + +@pytest.mark.parametrize( + "required_resource", + [ + ResourceIdentifier.from_string("images.segmentation_maps"), + ResourceIdentifier.from_string("psf_models.imaginary"), + ], +) +@patch("wf_psf.quality_control.pipeline.QualityControlConfigHandler") +@patch("wf_psf.quality_control.pipeline.QualityControlPipeline.validate_configuration") +def test_validate_metric_resources_unknown_resource( + mock_validate_configuration, + mock_config_handler, + qc_config_factory, + required_resource, +): + config = qc_config_factory(required_resources=[required_resource]) + mock_config_handler.return_value.load.return_value = config + mock_validate_configuration.return_value = None + + with pytest.raises( + ValueError, + match=( + f"Metric 'goodness_of_fit' requires unknown resource '{required_resource}'." + ), + ): + pipeline = QualityControlPipeline(qc_config_path="foo.yaml") + pipeline.validate_metric_resource_requirements() + + +@patch("wf_psf.quality_control.pipeline.QualityControlConfigHandler") +@patch("wf_psf.quality_control.pipeline.QualityControlPipeline.validate_configuration") +def test_validate_rejection_policy_metrics_all_valid( + mock_validate_configuration, mock_config_handler, qc_config_factory +): + mock_config_handler.return_value.load.return_value = qc_config_factory() + mock_validate_configuration.return_value = None + + pipeline = QualityControlPipeline(qc_config_path="foo.yaml") + with does_not_raise(): + pipeline.validate_rejection_policy_metrics() + + +@patch("wf_psf.quality_control.pipeline.QualityControlConfigHandler") +@patch("wf_psf.quality_control.pipeline.QualityControlPipeline.validate_configuration") +def test_validate_rejection_policy_disabled_policy( + mock_validate_configuration, mock_config_handler, qc_config_factory +): + # Define an invalid rejection policy that is disabled + rejection_policy = { + "foo_metric": RejectionPolicyConfig( + enabled=False, + diagnostic="bar", + policy={}, + ) + } + config = qc_config_factory(rejection=rejection_policy) + mock_config_handler.return_value.load.return_value = config + mock_validate_configuration.return_value = None + + # Verify no error is raised because disabled policies are skipped + pipeline = QualityControlPipeline(qc_config_path="foo.yaml") + pipeline.validate_rejection_policy_metrics() + + +@patch("wf_psf.quality_control.pipeline.QualityControlConfigHandler") +@patch("wf_psf.quality_control.pipeline.QualityControlPipeline.validate_configuration") +def test_validate_rejection_policy_metrics_unknown_metric( + mock_validate_configuration, mock_config_handler, qc_config_factory +): + rejection_policy = { + "foo_metric": RejectionPolicyConfig( + enabled=True, + diagnostic="bar", + policy={}, + ) + } + + config = qc_config_factory(rejection=rejection_policy) + mock_config_handler.return_value.load.return_value = config + mock_validate_configuration.return_value = None + + pipeline = QualityControlPipeline(qc_config_path="foo.yaml") + + with pytest.raises( + ValueError, + match="Rejection policy configured for unknown metric 'foo_metric'.", + ): + pipeline.validate_rejection_policy_metrics() + + +@patch("wf_psf.quality_control.pipeline.QualityControlConfigHandler") +@patch("wf_psf.quality_control.pipeline.QualityControlPipeline.validate_configuration") +def test_validate_rejection_policy_metrics_metric_not_enabled( + mock_validate_configuration, mock_config_handler, qc_config_factory +): + metric = { + "goodness_of_fit": QualityMetricConfig( + enabled=False, + required_resources=[], + ) + } + + config = qc_config_factory(metrics=metric) + mock_config_handler.return_value.load.return_value = config + mock_validate_configuration.return_value = None + + pipeline = QualityControlPipeline(qc_config_path="foo.yaml") + + with pytest.raises( + ValueError, + match="Rejection policy cannot be enabled because metric 'goodness_of_fit' is disabled.", + ): + pipeline.validate_rejection_policy_metrics() + + +@patch("wf_psf.quality_control.pipeline.QualityControlConfigHandler") +@patch("wf_psf.quality_control.pipeline.QualityControlPipeline.validate_configuration") +def test_validate_rejection_policy_metrics_metric_not_registered( + mock_validate_configuration, mock_config_handler, qc_config_factory +): + rejection_policy = { + "foo_metric": RejectionPolicyConfig( + enabled=True, + diagnostic="bar", + policy={}, + ) + } + + metric = { + "foo_metric": QualityMetricConfig( + enabled=True, + required_resources=[], + ) + } + config = qc_config_factory(metrics=metric, rejection=rejection_policy) + mock_config_handler.return_value.load.return_value = config + mock_validate_configuration.return_value = None + + pipeline = QualityControlPipeline(qc_config_path="foo.yaml") + + with pytest.raises( + KeyError, + match="Key 'foo_metric' not found.", + ): + pipeline.validate_rejection_policy_metrics() + + +@patch("wf_psf.quality_control.pipeline.QualityControlConfigHandler") +@patch("wf_psf.quality_control.pipeline.QualityControlPipeline.validate_configuration") +def test_validate_rejection_policy_metrics_diagnostic_not_found( + mock_validate_configuration, mock_config_handler, qc_config_factory +): + rejection_policy = { + "pixel_mask": RejectionPolicyConfig( + enabled=True, + diagnostic="bad_diagnostic", + policy={}, + ) + } + metric = { + "pixel_mask": QualityMetricConfig( + enabled=True, + required_resources=[], + ) + } + config = qc_config_factory(rejection=rejection_policy, metrics=metric) + mock_config_handler.return_value.load.return_value = config + mock_validate_configuration.return_value = None + pipeline = QualityControlPipeline(qc_config_path="foo.yaml") + + with pytest.raises( + ValueError, + match="Diagnostic 'bad_diagnostic' is not provided by the metric 'pixel_mask'.", + ): + pipeline.validate_rejection_policy_metrics() + + # Test pipeline runner def test_pipeline_run_single_rejection_policy(pipeline_factory): - metric_result = np.array([1.0, 2.0, 3.0]) + pixel_mask_metric_results = { + "total_masked_pixels": np.array([10.0, 20.0, 30.0]), + "total_masked_fraction": np.array([0.1, 0.2, 0.3]), + "aperture_masked_pixels": np.array([1.0, 2.0, 3.0]), + "aperture_masked_fraction": np.array([0.1, 0.2, 0.3]), + } + + gof_metric_results = { + "chi_square": np.array([123.0, 345.0, 678.0]), + "reduced_chi_square": np.array([1.2, 1.4, 1.1]), + } + validity_mask = np.array([True, False, True]) with ( patch.object( PixelMaskMetric, "compute", - return_value=metric_result, + return_value=pixel_mask_metric_results, ) as mock_mask_compute, patch.object( GoodnessOfFitMetric, "compute", - return_value=metric_result, + return_value=gof_metric_results, ) as mock_gof_compute, patch.object( ThresholdRejectionPolicy, @@ -96,17 +302,20 @@ def test_pipeline_run_single_rejection_policy(pipeline_factory): mock_mask_compute.assert_called_once() mock_gof_compute.assert_called_once() - mock_apply.assert_called_once_with(metric_result) - - assert np.array_equal( - result.metrics["pixel_mask"], - np.array([1.0, 2.0, 3.0]), + mock_apply.assert_called_once_with( + pixel_mask_metric_results["aperture_masked_fraction"] ) - assert np.array_equal( - result.metrics["goodness_of_fit"], - np.array([1.0, 2.0, 3.0]), - ) + for diagnostic_id, diagnostic_result in pixel_mask_metric_results.items(): + assert np.array_equal( + result.metrics["pixel_mask"][diagnostic_id], diagnostic_result + ) + + for diagnostic_id, diagnostic_result in gof_metric_results.items(): + assert np.array_equal( + result.metrics["goodness_of_fit"][diagnostic_id], + diagnostic_result, + ) assert np.array_equal( result.validity_masks["pixel_mask"], @@ -124,7 +333,17 @@ def test_pipeline_run_single_rejection_policy(pipeline_factory): def test_pipeline_run_multiple_rejection_policies(pipeline_factory): - metric_result = np.array([1.0, 2.0, 3.0]) + pixel_mask_metric_results = { + "total_masked_pixels": np.array([10.0, 20.0, 30.0]), + "total_masked_fraction": np.array([0.1, 0.2, 0.3]), + "aperture_masked_pixels": np.array([1.0, 2.0, 3.0]), + "aperture_masked_fraction": np.array([0.1, 0.2, 0.3]), + } + + gof_metric_results = { + "chi_square": np.array([123.0, 345.0, 678.0]), + "reduced_chi_square": np.array([1.2, 1.4, 1.1]), + } validity_masks = [ np.array([True, True, False]), np.array([True, False, True]), @@ -134,12 +353,12 @@ def test_pipeline_run_multiple_rejection_policies(pipeline_factory): patch.object( PixelMaskMetric, "compute", - return_value=metric_result, + return_value=pixel_mask_metric_results, ) as mock_mask_compute, patch.object( GoodnessOfFitMetric, "compute", - return_value=metric_result, + return_value=gof_metric_results, ) as mock_gof_compute, patch.object( ThresholdRejectionPolicy, @@ -161,17 +380,24 @@ def test_pipeline_run_multiple_rejection_policies(pipeline_factory): mock_mask_compute.assert_called_once() mock_gof_compute.assert_called_once() - assert mock_apply.call_count == 2 - - assert np.array_equal( - result.metrics["pixel_mask"], - np.array([1.0, 2.0, 3.0]), + mock_apply.assert_has_calls( + [ + call(pixel_mask_metric_results["aperture_masked_fraction"]), + call(gof_metric_results["reduced_chi_square"]), + ], + any_order=True, ) - assert np.array_equal( - result.metrics["goodness_of_fit"], - np.array([1.0, 2.0, 3.0]), - ) + for diagnostic_id, diagnostic_result in pixel_mask_metric_results.items(): + assert np.array_equal( + result.metrics["pixel_mask"][diagnostic_id], diagnostic_result + ) + + for diagnostic_id, diagnostic_result in gof_metric_results.items(): + assert np.array_equal( + result.metrics["goodness_of_fit"][diagnostic_id], + diagnostic_result, + ) assert np.array_equal( result.validity_masks["pixel_mask"], @@ -188,7 +414,71 @@ def test_pipeline_run_multiple_rejection_policies(pipeline_factory): np.array([True, False, False]), ) + +def test_pipeline_run_rejection_policy_disabled(pipeline_factory): + pixel_mask_metric_results = { + "total_masked_pixels": np.array([10.0, 20.0, 30.0]), + "total_masked_fraction": np.array([0.1, 0.2, 0.3]), + "aperture_masked_pixels": np.array([1.0, 2.0, 3.0]), + "aperture_masked_fraction": np.array([0.1, 0.2, 0.3]), + } + + gof_metric_results = { + "chi_square": np.array([123.0, 345.0, 678.0]), + "reduced_chi_square": np.array([1.2, 1.4, 1.1]), + } + + validity_mask = np.array([True, True, True]) + + with ( + patch.object( + PixelMaskMetric, + "compute", + return_value=pixel_mask_metric_results, + ) as mock_mask_compute, + patch.object( + GoodnessOfFitMetric, + "compute", + return_value=gof_metric_results, + ) as mock_gof_compute, + patch.object( + ThresholdRejectionPolicy, + "apply", + return_value=validity_mask, + ) as mock_apply, + ): + pipeline = pipeline_factory( + "valid/quality_control_rejection_policy_disabled.yaml" + ) + + dataset = np.array([1.0, 2.0, 3.0]) + provided_resources = {"psf_models.standard": np.array([1.0, 2.0, 3.0])} + + result = pipeline.run( + dataset=dataset, + provided_resources=provided_resources, + ) + + mock_mask_compute.assert_called_once() + mock_gof_compute.assert_called_once() + mock_apply.assert_not_called() + + for diagnostic_id, diagnostic_result in pixel_mask_metric_results.items(): + assert np.array_equal( + result.metrics["pixel_mask"][diagnostic_id], diagnostic_result + ) + + for diagnostic_id, diagnostic_result in gof_metric_results.items(): + assert np.array_equal( + result.metrics["goodness_of_fit"][diagnostic_id], + diagnostic_result, + ) + + assert "shapes" not in result.metrics + + assert result.validity_masks == {} + assert np.array_equal( result.valid_mask, - np.array([True, False, False]), + validity_mask, )