hackathon-ENGIN/scripts/build_model_artifact.py
Jakub Famulski 2 8e6dcb750d
All checks were successful
ENGIN CI / Build, test and smoke (push) Successful in 2m58s
prod: load versioned model artifact at runtime
2026-08-25 13:06:14 +02:00

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()