File size: 1,770 Bytes
0f16e07
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
import argparse
import json
from pathlib import Path

import joblib
import pandas as pd


MODEL_PATH = Path(__file__).parent / "model.joblib"

FEATURES = [
    "sector",
    "impact",
    "decision_autonomy",
    "human_oversight",
    "monitoring",
    "traceability",
    "technical_documentation",
]


def load_model():
    if not MODEL_PATH.exists():
        raise FileNotFoundError(
            f"Model artifact not found: {MODEL_PATH}"
        )

    return joblib.load(MODEL_PATH)


def validate_input(data: dict) -> None:
    missing = [feature for feature in FEATURES if feature not in data]

    if missing:
        raise ValueError(
            "Missing required features: " + ", ".join(missing)
        )


def predict_governance_risk(data: dict) -> dict:
    validate_input(data)

    model = load_model()

    frame = pd.DataFrame(
        [{feature: data[feature] for feature in FEATURES}]
    )

    prediction = model.predict(frame)[0]

    result = {
        "risk_tier": str(prediction),
    }

    if hasattr(model, "predict_proba"):
        probabilities = model.predict_proba(frame)[0]
        classes = model.classes_

        result["class_probabilities"] = {
            str(label): round(float(probability), 6)
            for label, probability in zip(classes, probabilities)
        }

    return result


def main():
    parser = argparse.ArgumentParser(
        description="Run the AIGov governance risk classifier."
    )

    parser.add_argument(
        "--json",
        required=True,
        help="Governance scenario encoded as JSON.",
    )

    args = parser.parse_args()

    data = json.loads(args.json)
    result = predict_governance_risk(data)

    print(json.dumps(result, indent=2))


if __name__ == "__main__":
    main()