"""Build or verify the versioned ENGIN production model artifact.""" from __future__ import annotations import argparse from pathlib import Path import pandas as pd from pandas.testing import assert_frame_equal from engin.artifact import ( LoadedModelArtifact, load_model_artifact, save_model_artifact, sha256_file, ) from engin.inference import predict_test, validate_submission from final_pipeline import ( LABEL_MODEL_NAME, SEVERITY_CANDIDATE_ID, train_models, ) DEFAULT_MODEL_VERSION = "engin-2026.08.25-1" def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--val", type=Path, default=Path("val.csv")) parser.add_argument("--test", type=Path, default=Path("test.csv")) parser.add_argument( "--expected-submission", type=Path, default=Path("predictions.csv") ) parser.add_argument( "--expected-diagnostics", type=Path, default=Path("prediction_diagnostics.csv"), ) parser.add_argument("--model-version", default=DEFAULT_MODEL_VERSION) parser.add_argument("--source-revision", default="unknown") parser.add_argument("--random-state", type=int, default=42) parser.add_argument("--n-jobs", type=int, default=-1) parser.add_argument("--artifact-dir", type=Path) parser.add_argument( "--verify-only", action="store_true", help="Load and verify the existing artifact without retraining.", ) return parser.parse_args() def _artifact_dir(args: argparse.Namespace) -> Path: return args.artifact_dir or Path("artifacts") / args.model_version def build_artifact(args: argparse.Namespace) -> LoadedModelArtifact: artifact_dir = _artifact_dir(args) if not args.verify_only: val = pd.read_csv(args.val).reset_index(drop=True) models = train_models( val, random_state=args.random_state, n_jobs=args.n_jobs, ) save_model_artifact( models, artifact_dir, model_version=args.model_version, source_revision=args.source_revision, training_data_sha256=sha256_file(args.val), random_state=args.random_state, n_jobs=args.n_jobs, label_model=LABEL_MODEL_NAME, severity_model=SEVERITY_CANDIDATE_ID, ) return load_model_artifact( artifact_dir, expected_model_version=args.model_version, ) def verify_predictions( artifact: LoadedModelArtifact, *, test_path: Path, expected_submission_path: Path, expected_diagnostics_path: Path, ) -> None: test = pd.read_csv(test_path).reset_index(drop=True) submission, diagnostics = predict_test(artifact.models, test) validate_submission(submission, test) expected_submission = pd.read_csv(expected_submission_path).reset_index(drop=True) expected_diagnostics = pd.read_csv(expected_diagnostics_path).reset_index(drop=True) assert_frame_equal( submission, expected_submission, check_dtype=False, check_exact=True, ) assert_frame_equal( diagnostics, expected_diagnostics, check_dtype=False, check_exact=False, rtol=1e-12, atol=1e-12, ) def main() -> None: args = parse_args() artifact = build_artifact(args) verify_predictions( artifact, test_path=args.test, expected_submission_path=args.expected_submission, expected_diagnostics_path=args.expected_diagnostics, ) metadata = artifact.metadata print(f"Artifact verified: {_artifact_dir(args).resolve()}") print(f"Model version: {metadata.model_version}") print(f"Model SHA-256: {metadata.model_sha256}") print(f"Training data SHA-256: {metadata.training_data_sha256}") print(f"Python: {metadata.python_version} | inference n_jobs: {metadata.n_jobs}") print("Predictions: exact 600-row regression match") if __name__ == "__main__": main()