All checks were successful
ENGIN CI / Build, test and smoke (push) Successful in 2m58s
131 lines
4.0 KiB
Python
131 lines
4.0 KiB
Python
"""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()
|