Coverage for packages/dqm-ml-core/src/dqm_ml_core/models/processors.py: 91%
146 statements
« prev ^ index » next coverage.py v7.14.1, created at 2026-07-21 08:27 +0000
« prev ^ index » next coverage.py v7.14.1, created at 2026-07-21 08:27 +0000
1"""Processor configuration models for DQM-ML pipelines.
3Defines configuration classes for all supported processor types:
4- Image feature extraction (luminosity, contrast, blur, entropy)
5- Neural network embedding extraction
6- Completeness metrics
7- Representativeness evaluation (chi-square, GRTE, KS, Shannon entropy)
8- Diversity metrics (Simpson, Gini, Shannon, richness)
9- Domain gap measurement (MMD, discriminative, etc.)
11Also includes supporting configuration for models, inference, kernels,
12distance metrics, and summary statistics.
13"""
15from typing import Annotated, Any, Literal
17from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
19from dqm_ml_core.models.columns import ColumnsConfig
21_LUMINOSITY_STANDARDS: dict[str, tuple[float, float, float]] = {
22 "bt601": (0.299, 0.587, 0.114),
23 "bt709": (0.2126, 0.7152, 0.0722),
24 "bt2020": (0.2627, 0.6780, 0.0593),
25}
28class _ProcessorBase(BaseModel):
29 """Base processor configuration with common fields.
31 Attributes:
32 name: Unique processor name.
33 type: Processor type discriminator for serialization.
34 columns: Column input/output configuration.
35 """
37 model_config = ConfigDict(extra="forbid")
39 name: str = Field(description="Unique processor name.")
40 type: str = Field(description="Processor type discriminator.")
41 columns: ColumnsConfig | None = None
44class HistogramConfig(BaseModel):
45 """Histogram parameters for feature extraction.
47 Attributes:
48 bins: Number of histogram bins (must be positive).
49 """
51 model_config = ConfigDict(extra="forbid")
53 bins: int = Field(default=256, gt=0, description="Number of histogram bins.")
56class ImageFeaturesProcessorConfig(_ProcessorBase):
57 """Configuration for low-level image feature extraction.
59 Extracts luminosity, contrast, blur, and entropy features from images.
61 Attributes:
62 type: Processor type discriminator ("image_features").
63 features: List of image features to compute.
64 batch_size: Batch size for image processing.
65 grayscale: Whether to convert images to grayscale.
66 normalize: Whether to normalize pixel values to [0, 1].
67 laplacian_kernel: Laplacian kernel size for blur detection ("3x3" or "5x5").
68 clip_percentiles: Percentile clipping for extreme pixel values, e.g. (1, 99).
69 histogram: Histogram configuration for feature computation.
70 luminosity_weights: Luminosity weights for grayscale conversion.
71 Standard name ('bt601', 'bt709', 'bt2020') or [R, G, B] list/tuple.
72 Defaults to BT.709 when None.
73 """
75 type: Literal["image_features"] = "image_features"
76 features: list[str] = Field(
77 default=["luminosity", "contrast", "blur", "entropy"],
78 description="List of image features to compute.",
79 )
80 batch_size: int = Field(default=64, gt=0, description="Batch size for image processing.")
81 grayscale: bool = Field(default=True, description="Convert images to grayscale.")
82 normalize: bool = Field(default=True, description="Normalise pixel values to [0, 1].")
83 laplacian_kernel: str = Field(default="3x3", description="Laplacian kernel size for blur detection.")
84 clip_percentiles: tuple[int, int] | None = Field(
85 default=None,
86 description="Percentile clipping for extreme pixel values, e.g. (1, 99).",
87 )
88 histogram: HistogramConfig | None = None
89 luminosity_weights: str | tuple[float, float, float] | None = Field(
90 default=None,
91 description="Luminosity weights for grayscale conversion. "
92 "Standard name ('bt601', 'bt709', 'bt2020') or [R, G, B] list. "
93 "Defaults to BT.709 when None.",
94 )
96 @field_validator("luminosity_weights", mode="before")
97 @classmethod
98 def _normalize_luminosity_weights(cls, v: Any) -> Any:
99 """Normalize luminosity weights input to standard key or tuple.
101 Args:
102 v: Input value - None, standard name string, or [R, G, B] list/tuple.
104 Returns:
105 Normalized value: None, standard key (e.g. "bt709"), or tuple of 3 floats.
107 Raises:
108 ValueError: If input is not a recognized standard, list/tuple of length 3, or None.
109 """
110 if v is None:
111 return v
112 if isinstance(v, str): 112 ↛ 113line 112 didn't jump to line 113 because the condition on line 112 was never true
113 key = v.lower().replace(".", "")
114 if key not in _LUMINOSITY_STANDARDS:
115 raise ValueError(f"Unknown luminosity standard '{v}'. Use one of {list(_LUMINOSITY_STANDARDS)}.")
116 return key
117 if isinstance(v, (list, tuple)): 117 ↛ 121line 117 didn't jump to line 121 because the condition on line 117 was always true
118 if len(v) != 3: 118 ↛ 119line 118 didn't jump to line 119 because the condition on line 118 was never true
119 raise ValueError(f"luminosity_weights must have exactly 3 elements, got {len(v)}.")
120 return tuple(v)
121 raise ValueError(f"luminosity_weights must be a standard name, [R,G,B] list, or None, got {type(v).__name__}.")
124class ModelConfig(BaseModel):
125 """Neural network model configuration for embedding extraction.
127 Attributes:
128 arch: Model architecture name (e.g., "resnet18", "resnet50").
129 n_layer_feature: Layer index (negative for reverse) or list of layer names
130 for feature extraction. Default -2 (second to last layer).
131 device: Device for model inference ("auto", "cpu", "cuda").
132 """
134 model_config = ConfigDict(extra="forbid")
136 arch: str = Field(default="resnet18", description="Model architecture name.")
137 n_layer_feature: int | list[str] = Field(
138 default=-2,
139 description="Layer index or list of layer names for feature extraction.",
140 )
141 device: Literal["auto", "cpu", "cuda"] = Field(
142 default="auto",
143 description="Device for model inference.",
144 )
147class InferConfig(BaseModel):
148 """Inference pre-processing settings for image embeddings.
150 Attributes:
151 batch_size: Inference batch size.
152 width: Resize width for input images.
153 height: Resize height for input images.
154 norm_mean: Per-channel mean used for normalisation (ImageNet defaults).
155 norm_std: Per-channel std used for normalisation (ImageNet defaults).
156 """
158 model_config = ConfigDict(extra="forbid")
160 batch_size: int = Field(default=32, gt=0, description="Inference batch size.")
161 width: int = Field(default=224, gt=0, description="Resize width for input images.")
162 height: int = Field(default=224, gt=0, description="Resize height for input images.")
163 norm_mean: list[float] = Field(
164 default=[0.485, 0.456, 0.406],
165 description="Per-channel mean used for normalisation.",
166 )
167 norm_std: list[float] = Field(
168 default=[0.229, 0.224, 0.225],
169 description="Per-channel std used for normalisation.",
170 )
173class FeaturesEmbeddingsProcessorConfig(_ProcessorBase):
174 """Configuration for neural-network embedding feature extraction.
176 Extracts deep learning embeddings from images using a configured model.
178 Attributes:
179 type: Processor type discriminator ("features_embeddings").
180 model: Neural network model configuration.
181 infer: Inference pre-processing settings.
182 """
184 type: Literal["features_embeddings"] = "features_embeddings"
185 model: ModelConfig = Field(default_factory=ModelConfig)
186 infer: InferConfig = Field(default_factory=InferConfig)
187 # luminosity_weights: TODO
190class CompletenessProcessorConfig(_ProcessorBase):
191 """Configuration for completeness metric computation.
193 Computes per-column and overall completeness (non-null) metrics.
195 Attributes:
196 type: Processor type discriminator ("completeness").
197 include_per_column: Include per-column completeness scores in output.
198 include_overall: Include overall completeness score in output.
199 include_metadata: Include metadata (total rows, null counts) in output.
200 """
202 type: Literal["completeness"] = "completeness"
203 include_per_column: bool = Field(default=True, description="Include per-column completeness scores.")
204 include_overall: bool = Field(default=True, description="Include overall completeness score.")
205 include_metadata: bool = Field(default=False, description="Include metadata in output.")
208class InterpretationConfig(BaseModel):
209 """Human-readable labels for representativeness and diversity results.
211 Attributes:
212 follows_distribution: Label when data follows the expected distribution.
213 does_not_follow_distribution: Label when data diverges from expected distribution.
214 high_diversity: Label for high diversity results.
215 low_diversity: Label for low diversity results.
216 high_representativeness: Label for high representativeness results.
217 low_representativeness: Label for low representativeness results.
218 """
220 model_config = ConfigDict(extra="forbid")
222 follows_distribution: str = Field(
223 default="fits target",
224 description="Label when data follows the distribution.",
225 )
226 does_not_follow_distribution: str = Field(default="diverges from target", description="Label when data diverges.")
227 high_diversity: str = Field(default="varied", description="Label for high diversity.")
228 low_diversity: str = Field(default="uniform", description="Label for low diversity.")
229 high_representativeness: str = Field(
230 default="representative",
231 description="Label for high representativeness.",
232 )
233 low_representativeness: str = Field(
234 default="under-represented",
235 description="Label for low representativeness.",
236 )
239class HistogramsConfig(BaseModel):
240 """Histogram parameters for representativeness evaluation.
242 Attributes:
243 bins: Number of histogram bins (must be positive).
244 """
246 model_config = ConfigDict(extra="forbid")
248 bins: int = Field(default=10, gt=0, description="Number of histogram bins.")
251class ShannonConfig(BaseModel):
252 """Shannon-entropy threshold configuration for representativeness.
254 Attributes:
255 threshold: Entropy threshold for flagging under-represented values.
256 """
258 model_config = ConfigDict(extra="forbid")
260 threshold: float = Field(default=2.0, description="Entropy threshold for flagging.")
263class GrteConfig(BaseModel):
264 """Gini-ratio-threshold-entropy (GRTE) configuration for representativeness.
266 Attributes:
267 threshold: Gini ratio threshold for flagging.
268 scaling_factor: Scaling applied to the Gini ratio (negative inverts).
269 """
271 model_config = ConfigDict(extra="forbid")
273 threshold: float = Field(default=0.5, description="Gini ratio threshold.")
274 scaling_factor: float = Field(default=-2.0, description="Scaling applied to the Gini ratio.")
277class KsConfig(BaseModel):
278 """Kolmogorov-Smirnov test configuration for representativeness.
280 Attributes:
281 sample_size: Number of samples for KS testing.
282 min_sample_size: Minimum samples required for KS test.
283 sample_divisor: Divisor for automatic sample-size calculation.
284 """
286 model_config = ConfigDict(extra="forbid")
288 sample_size: int = Field(default=500, gt=0, description="Number of samples for KS testing.")
289 min_sample_size: int = Field(default=50, gt=0, description="Minimum samples required for KS test.")
290 sample_divisor: int = Field(
291 default=20,
292 gt=0,
293 description="Divisor for automatic sample-size calculation.",
294 )
297class ColumnDistributionParams(BaseModel):
298 """Per-column distribution parameters for user_provided mean_std_estimation."""
300 column: str
301 mean: float | None = None
302 std: float | None = None
303 min: float | None = None
304 max: float | None = None
307class RepresentativenessProcessorConfig(_ProcessorBase):
308 """Configuration for representativeness evaluation against a reference distribution.
310 Evaluates how well a dataset represents a target distribution using
311 multiple statistical metrics.
313 Attributes:
314 type: Processor type discriminator ("representativeness").
315 metrics: List of representativeness metrics to compute.
316 alpha: Significance level for statistical tests.
317 epsilon: Small constant to avoid division by zero.
318 distribution: Expected reference distribution ("normal" or "uniform").
319 interpretation: Human-readable labels for results.
320 histogram: Histogram configuration for evaluation.
321 shannon: Shannon entropy threshold configuration.
322 grte: GRTE configuration.
323 ks: Kolmogorov-Smirnov test configuration.
324 """
326 type: Literal["representativeness"] = "representativeness"
327 metrics: list[str] = Field(
328 default=["chi-square", "grte", "kolmogorov-smirnov", "shannon-entropy"],
329 description="List of representativeness metrics to compute.",
330 )
331 alpha: float = Field(default=0.05, description="Significance level for statistical tests.")
332 epsilon: float = Field(
333 default=1e-9,
334 gt=0,
335 description="Small constant to avoid division by zero.",
336 )
337 distribution: Literal["normal", "uniform"] = Field(
338 default="normal",
339 description="Expected reference distribution.",
340 )
341 mean_std_estimation: Literal["from_first_batch", "per_batch", "from_all_data", "user_provided"] = Field(
342 default="from_first_batch",
343 description=(
344 "How distribution parameters (mean/std for normal, min/max for uniform) are estimated. "
345 "'from_first_batch' — estimate from the first batch and reuse (consistent with bin edges). "
346 "'per_batch' — re-estimate on each batch (use with high-variance data, risks insufficient_bins). "
347 "'user_provided' — use explicit parameters from distribution_params. "
348 "'from_all_data' — estimate from full dataset (not yet implemented)."
349 ),
350 )
351 expected_counts_method: Literal["cdf", "monte_carlo"] = Field(
352 default="cdf",
353 description=(
354 "Method for computing expected bin counts. "
355 "'cdf' — exact expected counts via CDF (deterministic). "
356 "'monte_carlo' — Monte Carlo sampling via RNG (stochastic)."
357 ),
358 )
359 distribution_params: list[ColumnDistributionParams] | None = Field(
360 default=None,
361 description=(
362 "Per-column explicit distribution parameters for 'user_provided' strategy. "
363 "Example: [{'column': 'col1', 'mean': 0.0, 'std': 1.0}]"
364 ),
365 )
366 interpretation: InterpretationConfig | None = None
367 histogram: HistogramsConfig | None = None
368 shannon: ShannonConfig | None = None
369 grte: GrteConfig | None = None
370 ks: KsConfig | None = None
373class DiversityProcessorConfig(_ProcessorBase):
374 """Configuration for diversity metric computation.
376 Computes diversity metrics: Simpson, Gini, Shannon entropy, and richness.
378 Attributes:
379 type: Processor type discriminator ("diversity").
380 metrics: List of diversity metrics to compute.
381 """
383 type: Literal["diversity"] = "diversity"
384 metrics: list[str] = Field(
385 default=["simpson", "gini", "shannon", "richness"],
386 description="List of diversity metrics to compute.",
387 )
390class KernelParamsRbf(BaseModel):
391 """RBF (Radial Basis Function) kernel parameters for distance metrics.
393 Attributes:
394 gamma: RBF kernel gamma parameter (inverse kernel width).
395 """
397 model_config = ConfigDict(extra="forbid")
399 gamma: float = Field(default=1.0, description="RBF kernel gamma parameter.")
402class KernelParamsPoly(BaseModel):
403 """Polynomial kernel parameters for distance metrics.
405 Attributes:
406 degree: Polynomial degree.
407 gamma: Polynomial kernel gamma parameter.
408 coefficient0: Polynomial kernel coefficient offset (constant term).
409 """
411 model_config = ConfigDict(extra="forbid")
413 degree: float = Field(default=3.0, description="Polynomial degree.")
414 gamma: float = Field(default=1.0, description="Polynomial kernel gamma.")
415 coefficient0: float = Field(default=1.0, description="Polynomial kernel coefficient offset.")
418class DistanceConfig(BaseModel):
419 """Distance metric configuration for domain-gap computation.
421 Attributes:
422 metric: Distance metric name (e.g., "mmd", "discriminative", "klmvn_diag").
423 evaluator: Optional evaluator type for discriminative metrics.
424 k: Number of nearest neighbours (for k-NN based metrics).
425 feature_weights: Per-feature weights for weighted distance computation.
426 kernel_params: Kernel parameters (RBF or Polynomial) for kernel-based metrics.
427 epsilon: Regularization epsilon for numerical stability of covariance-based metrics.
428 klmvn_var_eps: Variance regularization for KL divergence numerical stability.
429 """
431 model_config = ConfigDict(extra="forbid")
433 metric: str = Field(description="Distance metric name (e.g. 'mmd', 'discriminative').")
434 evaluator: str | None = Field(default=None, description="Optional evaluator type.")
435 k: int | None = Field(
436 default=None,
437 description="Number of nearest neighbours (if applicable).",
438 )
439 feature_weights: list[float] | None = Field(
440 default=None,
441 description="Per-feature weights for weighted distance computation.",
442 )
443 kernel_params: KernelParamsRbf | KernelParamsPoly | None = None
444 epsilon: float = Field(
445 default=1e-6,
446 ge=0,
447 description="Regularization epsilon for numerical stability of covariance-based metrics (e.g. FID).",
448 )
449 klmvn_var_eps: float = Field(
450 default=0.0,
451 ge=0,
452 description=(
453 "Variance regularization for klmvn_diag numerical stability "
454 "(default 0.0). When > 0, source and target variances are "
455 "replaced by var + klmvn_var_eps * mean(var) before computing "
456 "KL divergence, preventing blow-up from near-zero variance "
457 "dimensions."
458 ),
459 )
462class HistogramSummaryConfig(BaseModel):
463 """Histogram settings for embedding summary statistics.
465 Attributes:
466 dims: Number of dimensions to histogram.
467 bins: Number of bins per dimension (must be positive).
468 range: Histogram range [min, max] for each dimension.
469 """
471 model_config = ConfigDict(extra="forbid")
473 dims: int = Field(default=64, gt=0, description="Number of dimensions to histogram.")
474 bins: int = Field(default=32, gt=0, description="Number of bins per dimension.")
475 range: list[float] = Field(default=[-3.0, 3.0], description="Histogram range [min, max].")
478class SummaryConfig(BaseModel):
479 """Embedding summary configuration for domain-gap computation.
481 Attributes:
482 collect_sum_outer: Collect sum-of-outer-products for covariance estimation.
483 store_embeddings: Store full embedding vectors in output.
484 histogram: Histogram configuration for embedding summaries.
485 """
487 model_config = ConfigDict(extra="forbid")
489 collect_sum_outer: bool | None = Field(
490 default=None,
491 description="Collect sum-of-outer-products for covariance estimation.",
492 )
493 store_embeddings: bool | None = Field(default=None, description="Store full embedding vectors.")
494 histogram: HistogramSummaryConfig | None = None
497class DomainGapProcessorConfig(_ProcessorBase):
498 """Configuration for domain-gap (distribution shift) measurement.
500 Measures distribution shift between datasets using a configured distance metric.
502 Attributes:
503 type: Processor type discriminator ("domain_gap").
504 columns: Column input configuration (required).
505 distance: Distance metric configuration.
506 summary: Embedding summary configuration.
507 """
509 type: Literal["domain_gap"] = "domain_gap"
510 columns: ColumnsConfig = Field(description="Column input configuration (required).")
511 distance: DistanceConfig = Field(description="Distance metric configuration.")
512 summary: SummaryConfig | None = None
514 @model_validator(mode="after")
515 def _validate_domain_gap(self) -> "DomainGapProcessorConfig":
516 """Validate domain-gap configuration constraints."""
517 if self.columns and not self.columns.input: 517 ↛ 518line 517 didn't jump to line 518 because the condition on line 517 was never true
518 raise ValueError("'columns.input' is required and must not be empty for domain_gap processors")
519 if self.distance and self.distance.feature_weights is not None and self.columns and self.columns.input:
520 n_cols = len(self.columns.input)
521 if len(self.distance.feature_weights) != n_cols: 521 ↛ 522line 521 didn't jump to line 522 because the condition on line 521 was never true
522 raise ValueError(
523 f"feature_weights length ({len(self.distance.feature_weights)}) "
524 f"must match columns.input length ({n_cols})"
525 )
526 return self
529ProcessorConfig = Annotated[
530 ImageFeaturesProcessorConfig
531 | FeaturesEmbeddingsProcessorConfig
532 | CompletenessProcessorConfig
533 | RepresentativenessProcessorConfig
534 | DiversityProcessorConfig
535 | DomainGapProcessorConfig,
536 Field(discriminator="type"),
537]