Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion config/quality_control.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ metrics:
params:
aperture:
type: circular
centre: stamp_centre
radius: 2.6
unit: sigma

Expand All @@ -21,20 +22,21 @@ metrics:
required_resources:
- psf_models.standard
params:
statistic: reduced_chi_square
normalize_residuals: true


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
Expand Down
100 changes: 17 additions & 83 deletions src/wf_psf/quality_control/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)


Expand Down Expand Up @@ -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}' "
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -437,16 +375,16 @@ def load(self) -> QualityControlConfig:
Returns
-------
QualityControlConfig
Parsed and validated quality control configuration.
Parsed quality control configuration.

Raises
------
TypeError
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 = {}
Expand All @@ -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)
8 changes: 6 additions & 2 deletions src/wf_psf/quality_control/metrics/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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.

Expand All @@ -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
Expand Down
4 changes: 4 additions & 0 deletions src/wf_psf/quality_control/metrics/goodness_of_fit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down
9 changes: 9 additions & 0 deletions src/wf_psf/quality_control/metrics/pixel_masks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down
98 changes: 90 additions & 8 deletions src/wf_psf/quality_control/pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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]

Expand All @@ -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.
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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,
Expand Down
Loading
Loading