| name | ml-data-validation |
| description | Validate training and inference data for ML pipelines. Outputs schema contracts, statistical validation rules, drift detection, and data quality gates before model training. |
| argument-hint | ["model type","data sources","training frequency","feature count","known data issues"] |
| allowed-tools | Read, Write, Bash |
ML Data Validation
Garbage in, garbage out. ML data validation catches schema violations, statistical anomalies, and distribution drift before they corrupt model training or silently degrade inference. Without it, data issues manifest as mysterious model performance degradation weeks after the root cause.
Process
- Define the schema contract. Expected columns, types, value ranges, and cardinalities.
- Compute reference statistics. On a known-good training dataset: mean, std, percentiles, null rates, category distributions.
- Set validation rules. What deviations are acceptable? What trigger a warning vs a block?
- Validate at every pipeline stage. Raw data ingestion, feature engineering, training data split, inference input.
- Detect drift. Compare current data distribution against the reference. Alert when drift exceeds threshold.
- Gate on validation. Training blocked if critical validations fail.
- Log all validation results. Track data quality trends over time.
Great Expectations Suite
import great_expectations as ge
from great_expectations.core import ExpectationSuite
import pandas as pd
context = ge.get_context()
def build_expectations_from_data(df: pd.DataFrame, suite_name: str) -> ExpectationSuite:
validator = context.sources.pandas_default.read_dataframe(df)
validator.expectation_suite_name = suite_name
for col in df.columns:
null_rate = df[col].isna().mean()
if null_rate == 0:
validator.expect_column_values_to_not_be_null(col)
elif null_rate < 0.05:
validator.expect_column_values_to_not_be_null(col,
mostly=1 - null_rate * 1.5)
if df[col].dtype in ['int64', 'float64']:
q1, q99 = df[col].quantile([0.01, 0.99])
validator.expect_column_values_to_be_between(col,
min_value=float(q1 * 0.5 if q1 > 0 else q1 * 2),
max_value=float(q99 * 2 if q99 > 0 else q99 * 0.5),
mostly=0.99,
)
validator.expect_column_mean_to_be_between(col,
min_value=(df[col].mean() * ),
max_value=(df[col].mean() * ),
)
df[col].dtype == :
unique_count = df[col].nunique()
unique_count <= :
validator.expect_column_values_to_be_in_set(col,
value_set=(df[col].dropna().unique()))
validator.save_expectation_suite(discard_failed_expectations=)
validator.get_expectation_suite()
() -> :
validator = context.sources.pandas_default.read_dataframe(df)
results = validator.validate(expectation_suite_name=suite_name)
summary = {
: results.success,
: (results.results),
: ( r results.results r.success),
: [
r.expectation_config.expectation_type
r results.results
r.success r.expectation_config.meta.get() ==
],
}
summary
Custom Validation Rules
import pandas as pd
import numpy as np
from dataclasses import dataclass, field
from typing import Callable, List, Optional
@dataclass
class ValidationRule:
name: str
check: Callable[[pd.DataFrame], bool]
severity: str
message: str
@dataclass
class ValidationResult:
rule: str
passed: bool
severity: str
message: str
details: dict = field(default_factory=dict)
class MLDataValidator:
def __init__(self, reference_stats: dict = None):
self.reference = reference_stats or {}
self.rules: List[ValidationRule] = []
def add_rule(self, rule: ValidationRule):
self.rules.append(rule)
def validate(self, df: pd.DataFrame) -> List[ValidationResult]:
results = []
rule .rules:
:
passed = rule.check(df)
results.append(ValidationResult(
rule=rule.name, passed=passed,
severity=rule.severity, message=rule.message,
))
Exception e:
results.append(ValidationResult(
rule=rule.name, passed=,
severity=rule.severity,
message=,
))
results
() -> [, []]:
failures = [r.message r results r.passed r.severity == ]
(failures) == , failures
() -> MLDataValidator:
validator = MLDataValidator()
required_cols = [, , ,
, , ]
validator.add_rule(ValidationRule(
name=,
check= df: (c df.columns c required_cols),
severity=,
message=,
))
validator.add_rule(ValidationRule(
name=,
check= df: (df[].unique()).issubset({, , , }),
severity=,
message=,
))
validator.add_rule(ValidationRule(
name=,
check= df: (df) >= ,
severity=,
message=,
))
ref_churn_rate = reference_df[].mean()
validator.add_rule(ValidationRule(
name=,
check= df: (df[].mean() - ref_churn_rate) < ,
severity=,
message=,
))
post_churn_cols = [, , ]
validator.add_rule(ValidationRule(
name=,
check= df: (c df.columns c post_churn_cols),
severity=,
message=,
))
validator.add_rule(ValidationRule(
name=,
check= df: df[].duplicated().() == ,
severity=,
message=,
))
validator.add_rule(ValidationRule(
name=,
check= df: df[].between(, ).(),
severity=,
message=,
))
validator
Distribution Drift Detection
from scipy import stats
import numpy as np
class DriftDetector:
"""Detect statistical drift between reference and current datasets."""
def __init__(self, reference_df: pd.DataFrame):
self.reference = reference_df
def detect_drift(self, current_df: pd.DataFrame,
psi_threshold: float = 0.2,
ks_alpha: float = 0.05) -> dict:
drift_report = {}
for col in self.reference.columns:
if col not in current_df.columns:
drift_report[col] = {"status": "MISSING", "drift": True}
continue
if self.reference[col].dtype in ['int64', 'float64']:
stat, p_value = stats.ks_2samp(
self.reference[col].dropna(),
current_df[col].dropna()
)
psi = self._compute_psi(
self.reference[col].dropna(),
current_df[col].dropna()
)
drift_report[col] = {
: ,
: (stat, ),
: (p_value, ),
: (psi, ),
: p_value < ks_alpha psi > psi_threshold,
: psi > psi > ,
}
.reference[col].dtype == :
ref_dist = .reference[col].value_counts(normalize=)
cur_dist = current_df[col].value_counts(normalize=)
psi = ._compute_categorical_psi(ref_dist, cur_dist)
drift_report[col] = {
: ,
: (psi, ),
: ((current_df[col]) - (.reference[col])),
: psi > psi_threshold,
: psi > psi > ,
}
drifted = [col col, r drift_report.items() r.get()]
{
: drifted,
: (drifted) > ,
: [c c drifted
drift_report[c].get() == ],
: drift_report,
}
() -> :
breakpoints = np.percentile(reference, np.linspace(, , bins + ))
breakpoints = np.unique(breakpoints)
(breakpoints) < :
ref_counts, _ = np.histogram(reference, bins=breakpoints)
cur_counts, _ = np.histogram(current, bins=breakpoints)
ref_pct = (ref_counts + ) / (reference)
cur_pct = (cur_counts + ) / (current)
(np.((ref_pct - cur_pct) * np.log(ref_pct / cur_pct)))
() -> :
all_cats = (ref.index) | (cur.index)
psi =
cat all_cats:
r = ref.get(cat, )
c = cur.get(cat, )
psi += (r - c) * np.log(r / c)
psi
CI/CD Data Quality Gate
def run_training_pipeline(data_path: str, model_output: str):
df = pd.read_parquet(data_path)
reference = pd.read_parquet("s3://ml-data/reference/churn_train_v1.parquet")
validator = build_churn_model_validator(reference)
results = validator.validate(df)
can_proceed, failures = validator.gate(results)
if not can_proceed:
raise ValueError(f"Data validation failed — BLOCKING TRAINING:\n" +
"\n".join(failures))
warnings = [r.message for r in results if not r.passed and r.severity == "warning"]
if warnings:
print(f"WARNING: {len(warnings)} data quality warnings:\n" + "\n".join(warnings))
detector = DriftDetector(reference)
drift = detector.detect_drift(df)
if drift["critical_drift"]:
raise ValueError(f"Critical drift detected in features: {drift['critical_drift']}")
if drift["drift_detected"]:
print(f"WARNING: Drift detected in {len(drift['drifted_features'])} features")
print("✓ Data validation passed — proceeding with training")
train_model(df, model_output)
Anti-Patterns to Avoid
| Anti-Pattern | Problem | Fix |
|---|
| Validating only at training | Inference data can drift silently | Validate inference inputs in production |
| Hard-coded thresholds | Business context changes; thresholds become stale | Compute thresholds from reference data; review quarterly |
| Blocking on all warnings | Pipeline never runs | Distinguish critical (block) from warning (alert) |
| No reference dataset | Nothing to compare drift against | Snapshot the training dataset as the reference |
| Ignoring feature distribution | Training/serving skew degrades model | Compare feature distributions, not just schemas |
| Skipping validation on retrain | Assumes data quality is stable | Validate on every training run |
| Logging only failures | Can't track quality trends | Log all validation results; track quality scores over time |
10 Rules
- Validate data at every stage: ingestion, feature engineering, train/test split, and inference.
- Separate critical failures (block training) from warnings (alert team) — not everything needs to block.
- Compute validation thresholds from reference data — don't hard-code them.
- Drift detection requires a reference — snapshot the canonical training set on day one.
- PSI > 0.2 is a critical drift signal requiring investigation before retraining.
- Target leakage detection is mandatory — a high-performing model that leaks is worthless in production.
- Duplicate check before train/test split — duplicates across the split inflate test performance.
- Log all validation results with timestamps — you need the trend, not just pass/fail.
- Test the validator itself — validation code that doesn't catch real issues is false confidence.
- Validation gates are non-negotiable for production models — a pipeline that skips validation "just this once" will do it again.