diff --git a/config/quality_control.yaml b/config/quality_control.yaml index aac5553d..f11d04ca 100644 --- a/config/quality_control.yaml +++ b/config/quality_control.yaml @@ -18,7 +18,6 @@ metrics: enabled: true required_resources: - psf_models.standard - params: statistic: reduced_chi_square normalize_residuals: true @@ -28,7 +27,9 @@ rejection: mask_obscuration: enabled: true - threshold: 0.25 + policy: + threshold: + value: 0.25 goodness_of_fit: enabled: false diff --git a/src/wf_psf/quality_control/context.py b/src/wf_psf/quality_control/context.py new file mode 100644 index 00000000..f085de7f --- /dev/null +++ b/src/wf_psf/quality_control/context.py @@ -0,0 +1,31 @@ +"""Quality control context. + +Encapsulates contextual information, such as datasets and resolved resources, +required by the quality control pipeline and its metrics. + +:Authors: + Jennifer Pollack +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any + + +@dataclass +class QualityControlContext: + """Context shared across the quality control pipeline. + + Attributes + ---------- + dataset : Any + Dataset or data container supplied to the quality control pipeline. + + resources : dict[str, Any] + Ready-to-use resources required by enabled quality metrics, keyed by + resource identifier. + """ + + dataset: Any + resources: dict[str, Any] = field(default_factory=dict) diff --git a/src/wf_psf/quality_control/pipeline.py b/src/wf_psf/quality_control/pipeline.py index 0624e0ef..ad5ea163 100644 --- a/src/wf_psf/quality_control/pipeline.py +++ b/src/wf_psf/quality_control/pipeline.py @@ -2,11 +2,10 @@ Defines the orchestration layer for dataset quality control. -The QualityControlPipeline coordinates quality metric evaluation, -application of rejection policies, reporting, and dataset filtering. -Individual quality metrics and rejection policies are provided through -their respective interfaces, allowing new methods to be added without -modifying the pipeline implementation. +The QualityControlPipeline coordinates quality metric evaluation and +sample rejection. Individual quality metrics and rejection policies are +provided through their respective interfaces, allowing new methods to +be added without modifying the pipeline implementation. :Authors: Jennifer Pollack @@ -15,6 +14,19 @@ from dataclasses import dataclass import numpy as np +from wf_psf.quality_control.config import QualityControlConfigHandler +from wf_psf.quality_control.context import QualityControlContext +from wf_psf.quality_control.metrics.base import QualityMetric +from wf_psf.quality_control.metrics.registry import build_metrics_registry +from wf_psf.quality_control.rejection.base import RejectionPolicy +from wf_psf.quality_control.rejection.registry import build_rejection_policy_registry +from wf_psf.quality_control.resources import Resources + +import logging + +logger = logging.getLogger(__name__) + + @dataclass class QualityControlResult: """Results produced by the quality control pipeline. @@ -24,7 +36,7 @@ class QualityControlResult: metrics Computed quality metrics indexed by metric name. - rejection_masks + validity_masks Boolean validity masks produced by each rejection policy. valid_mask @@ -34,7 +46,7 @@ class QualityControlResult: metrics: dict[str, np.ndarray] - rejection_masks: dict[str, np.ndarray] + validity_masks: dict[str, np.ndarray] valid_mask: np.ndarray @@ -50,15 +62,129 @@ class QualityControlPipeline: of the pipeline execution. """ - def __init__( - self, - metrics, - rejection_policies, - ): - ... + 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() + + def _instantiate_metrics(self) -> dict[str, QualityMetric]: + """Instantiate enabled quality metric implementations from configuration. + + Returns + ------- + dict[str, QualityMetric] + Enabled quality metric implementations keyed by metric name. + + Notes + ----- + The quality control configuration is assumed to have been validated + before policy instantiation. + """ + metrics = {} + + for name, metric_config in self.config.metrics.items(): + if not metric_config.enabled: + logger.debug("Skipping metric %s: not enabled.", name) + continue + + metric_cls = self.metrics_registry.get(name) + + metrics[name] = metric_cls() + + logger.debug("Instantiated metrics: %s", list(metrics)) + + return metrics + + def _instantiate_rejection_policies(self) -> dict[str, RejectionPolicy]: + """Instantiate enabled rejection policy implementations. + + Returns + ------- + dict[str, RejectionPolicy] + Enabled rejection policy implementations keyed by metric name. + + Notes + ----- + The quality control configuration is assumed to have been validated + before policy instantiation. + """ + rejection_policies = {} + + for metric_name, rejection_config in self.config.rejection.items(): + if not rejection_config.enabled: + logger.debug("Skipping rejection policy %s: not enabled.", metric_name) + continue + + policy_name, policy_params = next(iter(rejection_config.policy.items())) + policy_cls = self.rejection_registry.get(policy_name) + + rejection_policies[metric_name] = policy_cls(**policy_params) + + logger.debug("Instantiated rejection policies: %s", list(rejection_policies)) + + return rejection_policies + + def _resolve_resources(self, provided_resources): + """Resolve resources required by enabled quality metrics. + + Parameters + ---------- + provided_resources : Mapping[str, Any] or None + Ready-to-use resources supplied by the pipeline caller, keyed by + resource identifier. + + Returns + ------- + dict[str, Any] + Resolved resources required by enabled quality metrics. + """ + resource_manager = Resources(self.config) + return resource_manager.resolve(provided_resources) + + def run(self, dataset, provided_resources=None): + """Run quality control pipeline. + + Parameters + ---------- + dataset : Any + Dataset or data container supplied to the quality control pipeline. + + provided_resources : Mapping[str, Any] or None + Ready-to-use resources supplied by the pipeline caller, keyed by resource identifier. + + Notes + ----- + The pipeline is expected to be invoked only when at least one quality metric is enabled in the quality control configuration. + """ + resolved_resources = self._resolve_resources( + provided_resources=provided_resources + ) + + context = QualityControlContext(dataset, resolved_resources) + + metrics = self._instantiate_metrics() + + metric_results = { + name: metric.compute(context) for name, metric in metrics.items() + } + + rejection_policies = self._instantiate_rejection_policies() - def run(self, dataset): + validity_masks = { + name: policy.apply(metric_results[name]) + for name, policy in rejection_policies.items() + } - ... + 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: + metric_result = next(iter(metric_results.values())) + valid_mask = np.ones(metric_result.shape, dtype=bool) - return QualityControlResult(...) \ No newline at end of file + return QualityControlResult( + metrics=metric_results, + validity_masks=validity_masks, + valid_mask=valid_mask, + ) diff --git a/src/wf_psf/quality_control/rejection/base.py b/src/wf_psf/quality_control/rejection/base.py index 38c5a87b..efcdfb28 100644 --- a/src/wf_psf/quality_control/rejection/base.py +++ b/src/wf_psf/quality_control/rejection/base.py @@ -37,7 +37,12 @@ class RejectionPolicy(ABC): @abstractmethod def apply(self, metric: np.ndarray) -> np.ndarray: - """Return a boolean mask identifying valid dataset samples. - - The returned mask has one entry per dataset sample. + """Apply the rejection policy to metric values. + + Returns + ------- + np.ndarray + Boolean validity mask with one entry per dataset sample. ``True`` + indicates that the sample passes the rejection policy and should + be retained; ``False`` indicates that it should be rejected. """ diff --git a/src/wf_psf/quality_control/rejection/threshold.py b/src/wf_psf/quality_control/rejection/threshold.py index f4936059..3a2590b5 100644 --- a/src/wf_psf/quality_control/rejection/threshold.py +++ b/src/wf_psf/quality_control/rejection/threshold.py @@ -16,10 +16,23 @@ class ThresholdRejectionPolicy(RejectionPolicy): - """Reject dataset samples based on configurable metric thresholds.""" + """Reject dataset samples based on configurable metric threshold. + + Attributes + ---------- + name : str + Policy identifier used by the rejection policy registry. + + value : float + Threshold applied by the rejection policy. + + """ name = "threshold" + def __init__(self, value: float): + self.value = value + def apply(self, metric: np.ndarray) -> np.ndarray: """Apply threshold-based rejection to metric values.""" raise NotImplementedError 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 index 0dc25183..1b76b0fe 100644 --- 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 @@ -19,4 +19,6 @@ rejection: mask_obscuration: enabled: true - threshold: 0.25 + policy: + threshold: + value: 0.25 diff --git a/src/wf_psf/tests/test_quality_control/data/valid/imaginary_metric.yaml b/src/wf_psf/tests/test_quality_control/data/valid/imaginary_metric.yaml new file mode 100644 index 00000000..eb6ebf8d --- /dev/null +++ b/src/wf_psf/tests/test_quality_control/data/valid/imaginary_metric.yaml @@ -0,0 +1,4 @@ +metrics: + + imaginary_metric: + enabled: true \ No newline at end of file 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 d7222a92..7a5e4636 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 @@ -18,7 +18,6 @@ metrics: enabled: true required_resources: - psf_models.standard - params: statistic: reduced_chi_square normalize_residuals: true 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 new file mode 100644 index 00000000..d94f7e76 --- /dev/null +++ b/src/wf_psf/tests/test_quality_control/data/valid/quality_control_multiple_rejection_policies.yaml @@ -0,0 +1,42 @@ +resources: + + psf_models: + standard: + inference_config: inference_standard.yaml + oversampled: + inference_config: inference_oversampled.yaml + +metrics: + + mask_obscuration: + enabled: true + params: + aperture: gaussian + sigma: 2.5 + + goodness_of_fit: + enabled: true + required_resources: + - psf_models.standard + params: + statistic: reduced_chi_square + normalize_residuals: true + +rejection: + + mask_obscuration: + enabled: true + policy: + threshold: + value: 3.0 + + goodness_of_fit: + enabled: true + policy: + threshold: + value: 3.0 + +reporting: + + save_metrics: true + log_statistics: true \ No newline at end of file diff --git a/src/wf_psf/tests/test_quality_control/pipeline_test.py b/src/wf_psf/tests/test_quality_control/pipeline_test.py new file mode 100644 index 00000000..d3cbdd88 --- /dev/null +++ b/src/wf_psf/tests/test_quality_control/pipeline_test.py @@ -0,0 +1,194 @@ +import numpy as np +from pathlib import Path +import pytest +from unittest.mock import patch + +from wf_psf.quality_control.pipeline import QualityControlPipeline +from wf_psf.quality_control.config import QualityControlConfig +from wf_psf.quality_control.metrics.mask_obscuration import MaskObscurationMetric +from wf_psf.quality_control.metrics.goodness_of_fit import GoodnessOfFitMetric +from wf_psf.quality_control.rejection.threshold import ThresholdRejectionPolicy + + +@pytest.fixture +def pipeline_factory(): + def build(config_file): + path = Path(__file__).parent / "data" / config_file + return QualityControlPipeline(path) + + return build + + +def test_pipeline_constructor(pipeline_factory): + pipeline = pipeline_factory("valid/quality_control.yaml") + + # Check config + assert isinstance(pipeline.config, QualityControlConfig) + + # Check metrics registry + assert pipeline.metrics_registry.get("mask_obscuration") is MaskObscurationMetric + assert pipeline.metrics_registry.get("goodness_of_fit") is GoodnessOfFitMetric + + # Check rejection registry + assert pipeline.rejection_registry.get("threshold") is ThresholdRejectionPolicy + + +def test_pipeline_instantiate_metrics_valid(pipeline_factory): + pipeline = pipeline_factory("valid/quality_control.yaml") + + metrics = pipeline._instantiate_metrics() + + assert len(metrics) == 2 + assert isinstance(metrics["mask_obscuration"], MaskObscurationMetric) + assert isinstance(metrics["goodness_of_fit"], GoodnessOfFitMetric) + + +def test_pipeline_instantiate_metrics_unknown_metric(pipeline_factory): + pipeline = pipeline_factory("valid/imaginary_metric.yaml") + + with pytest.raises(KeyError): + pipeline._instantiate_metrics() + + +def test_pipeline_instantiate_rejection_policy_valid(pipeline_factory): + pipeline = pipeline_factory("valid/quality_control.yaml") + + rejection_policies = pipeline._instantiate_rejection_policies() + + assert len(rejection_policies) == 1 + assert isinstance(rejection_policies["mask_obscuration"], ThresholdRejectionPolicy) + assert rejection_policies["mask_obscuration"].value == 3.0 + + assert "goodness_of_fit" not in rejection_policies + + +# Test pipeline runner +def test_pipeline_run_single_rejection_policy(pipeline_factory): + metric_result = np.array([1.0, 2.0, 3.0]) + validity_mask = np.array([True, False, True]) + + with ( + patch.object( + MaskObscurationMetric, + "compute", + return_value=metric_result, + ) as mock_mask_compute, + patch.object( + GoodnessOfFitMetric, + "compute", + return_value=metric_result, + ) as mock_gof_compute, + patch.object( + ThresholdRejectionPolicy, + "apply", + return_value=validity_mask, + ) as mock_apply, + ): + pipeline = pipeline_factory("valid/quality_control.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_called_once_with(metric_result) + + assert np.array_equal( + result.metrics["mask_obscuration"], + np.array([1.0, 2.0, 3.0]), + ) + + assert np.array_equal( + result.metrics["goodness_of_fit"], + np.array([1.0, 2.0, 3.0]), + ) + + assert np.array_equal( + result.validity_masks["mask_obscuration"], + np.array([True, False, True]), + ) + + assert "goodness_of_fit" not in result.validity_masks + assert "shapes" not in result.metrics + assert "shapes" not in result.validity_masks + + assert np.array_equal( + result.valid_mask, + np.array([True, False, True]), + ) + + +def test_pipeline_run_multiple_rejection_policies(pipeline_factory): + metric_result = np.array([1.0, 2.0, 3.0]) + validity_masks = [ + np.array([True, True, False]), + np.array([True, False, True]), + ] + + with ( + patch.object( + MaskObscurationMetric, + "compute", + return_value=metric_result, + ) as mock_mask_compute, + patch.object( + GoodnessOfFitMetric, + "compute", + return_value=metric_result, + ) as mock_gof_compute, + patch.object( + ThresholdRejectionPolicy, + "apply", + side_effect=validity_masks, + ) as mock_apply, + ): + pipeline = pipeline_factory( + "valid/quality_control_multiple_rejection_policies.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() + assert mock_apply.call_count == 2 + + assert np.array_equal( + result.metrics["mask_obscuration"], + np.array([1.0, 2.0, 3.0]), + ) + + assert np.array_equal( + result.metrics["goodness_of_fit"], + np.array([1.0, 2.0, 3.0]), + ) + + assert np.array_equal( + result.validity_masks["mask_obscuration"], + np.array([True, True, False]), + ) + + assert np.array_equal( + result.validity_masks["goodness_of_fit"], + np.array([True, False, True]), + ) + + assert np.array_equal( + result.valid_mask, + np.array([True, False, False]), + ) + + assert np.array_equal( + result.valid_mask, + np.array([True, False, False]), + )