josibra commited on
Commit
ca542ab
·
1 Parent(s): e6030d4

Déploiement auto depuis GitHub Actions avec LFS

Browse files
.gitattributes CHANGED
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ *.cbm filter=lfs diff=lfs merge=lfs -text
Dockerfile ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM python:3.12-slim
2
+
3
+ # 1. Création d'un utilisateur standard (UID 1000 requis par Hugging Face)
4
+ RUN useradd -m -u 1000 user
5
+ USER user
6
+
7
+ # 2. Variables d'environnement pour l'utilisateur et Python
8
+ ENV HOME=/home/user \
9
+ PATH=/home/user/.local/bin:$PATH \
10
+ PYTHONDONTWRITEBYTECODE=1 \
11
+ PYTHONUNBUFFERED=1
12
+
13
+ # 3. Répertoire de travail
14
+ WORKDIR $HOME/app
15
+
16
+ # 4. Copie des dépendances avec les bons droits
17
+ COPY --chown=user requirements.txt .
18
+
19
+ # 5. Mise à jour de pip (pour enlever le 'notice') et installation
20
+ RUN pip install --no-cache-dir --upgrade pip && \
21
+ pip install --no-cache-dir -r requirements.txt
22
+
23
+ # 6. Copie du code et des modèles
24
+ COPY --chown=user app/ ./app/
25
+ COPY --chown=user models/ ./models/
26
+
27
+ # 7. Variables de l'API avec le nouveau chemin
28
+ ENV GLOBAL_THRESHOLD=0.45
29
+ ENV MODEL_PATH=$HOME/app/models/pmvl_catboost_final.cbm
30
+ ENV MODEL_FEATURES_PATH=$HOME/app/models/pmvl_feature_columns.txt
31
+
32
+ # 8. Exposer le port 7860 (Port par défaut de Hugging Face Spaces)
33
+ EXPOSE 7860
34
+
35
+ # 9. Démarrage de l'API sur le port 7860
36
+ CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "7860"]
app/__init__.py ADDED
File without changes
app/main.py ADDED
@@ -0,0 +1,128 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from fastapi import FastAPI, HTTPException
2
+ from fastapi.middleware.cors import CORSMiddleware
3
+ import os
4
+ import pandas as pd
5
+ import numpy as np
6
+
7
+ from .schemas import PMVLFeatures, PredictionResponse
8
+ from .model_loader import get_model, get_feature_columns
9
+
10
+ GLOBAL_THRESHOLD_ENV = "GLOBAL_THRESHOLD"
11
+ DEFAULT_THRESHOLD = 0.45
12
+
13
+ app = FastAPI(
14
+ title="API de prédiction de la qualité PMVL",
15
+ description="API exposant le modèle CatBoost pour estimer la précision des PMVL.",
16
+ version="1.0.0",
17
+ )
18
+
19
+ app.add_middleware(
20
+ CORSMiddleware,
21
+ allow_origins=["*"],
22
+ allow_credentials=True,
23
+ allow_methods=["*"],
24
+ allow_headers=["*"],
25
+ )
26
+
27
+ GROUP_KEYS = [
28
+ "PMVL[ENTITE]",
29
+ "PMVL[Selected Fund code]",
30
+ "PMVL[ISIN]",
31
+ "PMVL[Ref Unik Asset]",
32
+ ]
33
+
34
+ def make_position_group(df: pd.DataFrame) -> pd.Series:
35
+ """Recrée la colonne position_group comme dans le notebook"""
36
+ existing = [c for c in GROUP_KEYS if c in df.columns]
37
+ if not existing:
38
+ return pd.Series(
39
+ np.arange(len(df)).astype(str), index=df.index, name="position_group"
40
+ )
41
+ return (
42
+ df[existing]
43
+ .astype(str)
44
+ .fillna("NA")
45
+ .agg("||".join, axis=1)
46
+ .rename("position_group")
47
+ )
48
+
49
+ def prepare_catboost_features(X: pd.DataFrame, cat_cols: list):
50
+ """Prépare les features (remplace None par 'MISSING' pour les cat, etc.)"""
51
+ X = X.copy()
52
+ for c in X.columns:
53
+ if c in cat_cols:
54
+ # Remplacement des valeurs manquantes par 'MISSING' et conversion en string
55
+ X[c] = X[c].fillna("MISSING").astype(str)
56
+ else:
57
+ # Remplacement par NaN et forçage numérique
58
+ X[c] = pd.to_numeric(X[c], errors="coerce")
59
+
60
+ for c in X.columns:
61
+ if X[c].dtype == bool:
62
+ X[c] = X[c].astype(int)
63
+ return X
64
+
65
+
66
+ @app.on_event("startup")
67
+ def load_model_on_startup():
68
+ try:
69
+ get_model()
70
+ get_feature_columns()
71
+ except Exception as e:
72
+ print(f"Erreur critique lors du chargement initial : {e}")
73
+
74
+ @app.get("/health", tags=["diagnostic"])
75
+ def health_check():
76
+ return {"status": "ok", "message": "L'API PMVL est opérationnelle."}
77
+
78
+ @app.post("/predict", response_model=PredictionResponse, tags=["prédiction"])
79
+ def predict_pmvl(features: PMVLFeatures):
80
+ try:
81
+ model = get_model()
82
+ feature_columns = get_feature_columns()
83
+
84
+ # 1) Récupérer les données brutes
85
+ raw_dict = features.model_dump(by_alias=True)
86
+ raw_dict.pop("PMVL[Holding date]", None)
87
+
88
+ # 2) Construire un DataFrame avec TOUTES les colonnes attendues
89
+ row = {col: np.nan for col in feature_columns}
90
+ for k, v in raw_dict.items():
91
+ if k in row and v is not None:
92
+ row[k] = v
93
+
94
+ raw_df = pd.DataFrame([row])
95
+
96
+ # 3) Recréer position_group si elle fait partie des features
97
+ if "position_group" in feature_columns:
98
+ raw_df["position_group"] = make_position_group(raw_df)
99
+
100
+ # 4) S'assurer de l'ordre exact des colonnes
101
+ df_input = raw_df[feature_columns]
102
+
103
+ # 5) Appliquer le nettoyage en demandant AU MODÈLE quelles sont les colonnes catégorielles !
104
+ cat_indices = model.get_cat_feature_indices()
105
+ cat_cols = [feature_columns[i] for i in cat_indices]
106
+
107
+ df_input = prepare_catboost_features(df_input, cat_cols)
108
+
109
+ # 6) Prédire
110
+ proba = float(model.predict_proba(df_input)[:, 1][0])
111
+
112
+ threshold = float(os.getenv(GLOBAL_THRESHOLD_ENV, DEFAULT_THRESHOLD))
113
+ prediction_bool = bool(proba >= threshold)
114
+
115
+ return PredictionResponse(
116
+ proba_bonne_estimation=proba,
117
+ prediction=prediction_bool,
118
+ seuil_applique=threshold,
119
+ fund_code=features.fund_code,
120
+ ref_unik_asset=features.ref_unik_asset,
121
+ )
122
+
123
+
124
+ except Exception as e:
125
+ raise HTTPException(
126
+ status_code=500,
127
+ detail=f"Erreur lors du traitement de la prédiction : {e}",
128
+ )
app/model_loader.py ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ from functools import lru_cache
3
+ from catboost import CatBoostClassifier
4
+
5
+ # Variables d'environnement pour rendre le chemin dynamique (utile pour Docker)
6
+ MODEL_PATH_ENV = "MODEL_PATH"
7
+ DEFAULT_MODEL_PATH = "models/pmvl_catboost_final.cbm"
8
+
9
+ FEATURES_PATH_ENV = "FEATURES_PATH"
10
+ DEFAULT_FEATURES_PATH = "models/pmvl_feature_columns.txt"
11
+
12
+ @lru_cache(maxsize=1)
13
+ def get_model() -> CatBoostClassifier:
14
+ """
15
+ Charge le modèle CatBoost depuis le disque.
16
+ Grâce à @lru_cache, cette fonction ne s'exécute réellement qu'une seule fois.
17
+ """
18
+ model_path = os.getenv(MODEL_PATH_ENV, DEFAULT_MODEL_PATH)
19
+
20
+ if not os.path.exists(model_path):
21
+ raise FileNotFoundError(f"Le fichier du modèle est introuvable au chemin : {model_path}")
22
+
23
+ print("Chargement du modèle CatBoost en mémoire...")
24
+ model = CatBoostClassifier()
25
+ model.load_model(model_path)
26
+ print("Modèle chargé avec succès !")
27
+
28
+ return model
29
+
30
+ #chargeur pour la liste de colonnes utilisées par le modèle (pour s'assurer que l'ordre des features est correct)
31
+ @lru_cache(maxsize=1)
32
+ def get_feature_columns() -> list[str]:
33
+ features_path = os.getenv(FEATURES_PATH_ENV, DEFAULT_FEATURES_PATH)
34
+ if not os.path.exists(features_path):
35
+ raise FileNotFoundError(f"Le fichier des features est introuvable au chemin : {features_path}")
36
+ with open(features_path, encoding="utf-8") as f:
37
+ cols = [line.strip() for line in f if line.strip()]
38
+ return cols
app/schemas.py ADDED
@@ -0,0 +1,153 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from pydantic import BaseModel, Field, validator
2
+ from typing import Optional
3
+ from datetime import date
4
+
5
+ class PMVLFeatures(BaseModel):
6
+ """
7
+ Schéma Pydantic représentant une observation d'entrée pour la prédiction PMVL.
8
+ Les noms des attributs sont des alias Python-friendly pour les colonnes complexes du DataFrame.
9
+ """
10
+ # ------------------ Colonnes de type Date ------------------
11
+ holding_date: date = Field(
12
+ ...,
13
+ alias="PMVL[Holding date]",
14
+ description="Date de la position (Trading/Holding date)"
15
+ )
16
+ quote_date: Optional[str] = Field(
17
+ None,
18
+ alias="PMVL[Quote Date]",
19
+ description="Date de la cotation (peut contenir des valeurs manquantes)"
20
+ )
21
+
22
+ # ------------------ Colonnes Numériques (float64) ------------------
23
+ pmvl_estim: Optional[float] = Field(
24
+ None,
25
+ alias="PMVL[PMVL Estimé]",
26
+ description="Estimation de la PMVL au temps t"
27
+ )
28
+ pfm_indice_perf: Optional[float] = Field(
29
+ None,
30
+ alias="PMVL[PFMIndice.Perf J / J-1]",
31
+ description="Performance de l'indice J / J-1"
32
+ )
33
+ prmp_pmvl: Optional[float] = Field(
34
+ None,
35
+ alias="PMVL[PRMP PMVL]",
36
+ description="Realized PMVL at t"
37
+ )
38
+ prmp_vnc: Optional[float] = Field(
39
+ None,
40
+ alias="PMVL[PRMP VNC]",
41
+ description="Realized VNC at t"
42
+ )
43
+ prmp_mtm: Optional[float] = Field(
44
+ None,
45
+ alias="PMVL[PRMP MtM]",
46
+ description="Realized MtM at t"
47
+ )
48
+ quantity: float = Field(
49
+ ...,
50
+ alias="PMVL[Quantity]",
51
+ description="Quantité de l'actif"
52
+ )
53
+ purch_val_clean: float = Field(
54
+ ...,
55
+ alias="PMVL[Purch. Val. (clean) (ptf cur.)]",
56
+ description="Valeur d'achat clean en devise du portefeuille"
57
+ )
58
+ quote: float = Field(
59
+ ...,
60
+ alias="PMVL[Quote]",
61
+ description="Cotation de l'actif"
62
+ )
63
+ vnc_agrege_dirty: float = Field(
64
+ ...,
65
+ alias="PMVL[VNC Agrege dirty (ptf cur.)]",
66
+ description="VNC Agrégée dirty en devise du portefeuille"
67
+ )
68
+
69
+ # ------------------ Colonnes Catégorielles (object) ------------------
70
+ entite: str = Field(
71
+ ...,
72
+ alias="PMVL[ENTITE]",
73
+ description="Entité légale"
74
+ )
75
+ isin: str = Field(
76
+ ...,
77
+ alias="PMVL[ISIN]",
78
+ description="Code ISIN de l'instrument"
79
+ )
80
+ orig_name: str = Field(
81
+ ...,
82
+ alias="PMVL[Orig. name]",
83
+ description="Nom original de l'actif"
84
+ )
85
+ ticker: str = Field(
86
+ ...,
87
+ alias="PMVL[Parametres_Indices.TICKER]",
88
+ description="Ticker de l'indice"
89
+ )
90
+ ref_unik_asset: str = Field(
91
+ ...,
92
+ alias="PMVL[Ref Unik Asset]",
93
+ description="Référence unique de l'asset"
94
+ )
95
+ fund_code: str = Field(
96
+ ...,
97
+ alias="PMVL[Selected Fund code]",
98
+ description="Code du fonds sélectionné"
99
+ )
100
+ col_3a: str = Field(
101
+ ...,
102
+ alias="PMVL[3A]",
103
+ description="Classification 3A"
104
+ )
105
+ canton: str = Field(
106
+ ...,
107
+ alias="PMVL[CANTON]",
108
+ description="Canton"
109
+ )
110
+ cic: str = Field(
111
+ ...,
112
+ alias="PMVL[CIC]",
113
+ description="Classification CIC"
114
+ )
115
+ groupe: str = Field(
116
+ ...,
117
+ alias="PMVL[GROUPE]",
118
+ description="Groupe"
119
+ )
120
+ ptf_name: str = Field(
121
+ ...,
122
+ alias="PMVL[Ptf name]",
123
+ description="Nom du portefeuille"
124
+ )
125
+
126
+ class Config:
127
+ populate_by_name = True # Nouvelle syntaxe Pydantic V2
128
+
129
+
130
+ class PredictionResponse(BaseModel):
131
+ """
132
+ Schéma de la réponse retournée par l'API.
133
+ """
134
+ proba_bonne_estimation: float = Field(
135
+ ...,
136
+ description="Probabilité (entre 0 et 1) que la PMVL soit une bonne estimation."
137
+ )
138
+ prediction: bool = Field(
139
+ ...,
140
+ description="True si la PMVL est jugée bonne (proba >= seuil), False sinon."
141
+ )
142
+ seuil_applique: float = Field(
143
+ ...,
144
+ description="Le seuil de probabilité utilisé pour cette décision (ex: 0.45)."
145
+ )
146
+ fund_code: str = Field(
147
+ ...,
148
+ description="Rappel du code du fonds pour traçabilité."
149
+ )
150
+ ref_unik_asset: str = Field(
151
+ ...,
152
+ description="Rappel de la référence de l'actif pour traçabilité."
153
+ )
models/pmvl_catboost_final.cbm ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b81635e7410aee5b4d95784954ea3e6d9356bcfbfdce88f713aa9086bb31fc14
3
+ size 6085568
models/pmvl_feature_columns.txt ADDED
@@ -0,0 +1,22 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ PMVL[ENTITE]
2
+ PMVL[ISIN]
3
+ PMVL[Orig. name]
4
+ PMVL[Parametres_Indices.TICKER]
5
+ PMVL[PFMIndice.Perf J / J-1]
6
+ PMVL[PRMP PMVL]
7
+ PMVL[PRMP VNC]
8
+ PMVL[PRMP MtM]
9
+ PMVL[Ref Unik Asset]
10
+ PMVL[Selected Fund code]
11
+ PMVL[3A]
12
+ PMVL[CANTON]
13
+ PMVL[CIC]
14
+ PMVL[GROUPE]
15
+ PMVL[Ptf name]
16
+ PMVL[Quantity]
17
+ PMVL[Purch. Val. (clean) (ptf cur.)]
18
+ PMVL[Quote Date]
19
+ PMVL[Quote]
20
+ PMVL[VNC Agrege dirty (ptf cur.)]
21
+ PMVL[PMVL Estimé]
22
+ position_group
requirements.txt ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ fastapi>=0.111.0
2
+ uvicorn[standard]>=0.30.1
3
+ pydantic>=2.7.0
4
+ catboost==1.2.10
5
+ pandas>=2.2.0
6
+ numpy>=1.26.0
7
+ pytest>=8.0.0
8
+ httpx>=0.27.0
tests/__init__.py ADDED
File without changes
tests/test_api.py ADDED
@@ -0,0 +1,85 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from fastapi.testclient import TestClient
2
+ from app.main import app
3
+
4
+ # Création du client de test FastAPI
5
+ client = TestClient(app)
6
+
7
+ # Payload valide (contenant tous les champs obligatoires définis dans schemas.py)
8
+ # On utilise ici les noms des variables Python (grâce à populate_by_name=True)
9
+ VALID_PAYLOAD = {
10
+ "holding_date": "2026-03-01",
11
+ "pmvl_estim": 1500.50,
12
+ "quantity": 100.0,
13
+ "purch_val_clean": 45000.0,
14
+ "quote": 510.0,
15
+ "vnc_agrege_dirty": 49000.0,
16
+ "entite": "ENTITE_TEST",
17
+ "isin": "FR0000000001",
18
+ "orig_name": "Asset Name Test",
19
+ "ticker": "TICKER_TEST",
20
+ "ref_unik_asset": "REF_12345",
21
+ "fund_code": "FUND_001",
22
+ "col_3a": "3A_TEST",
23
+ "canton": "CANTON_TEST",
24
+ "cic": "CIC_TEST",
25
+ "groupe": "GROUPE_TEST",
26
+ "ptf_name": "PTF_TEST"
27
+ }
28
+
29
+
30
+ def test_health_check():
31
+ """
32
+ Test 1 : Vérifie que le endpoint de diagnostic répond bien 200 OK.
33
+ """
34
+ response = client.get("/health")
35
+ assert response.status_code == 200
36
+ assert response.json()["status"] == "ok"
37
+
38
+
39
+ def test_predict_valid_data():
40
+ """
41
+ Test 2 : Vérifie qu'un payload complet et valide renvoie bien une prédiction.
42
+ """
43
+ response = client.post("/predict", json=VALID_PAYLOAD)
44
+
45
+ # Ligne ajoutée pour debug
46
+ print("RESPONSE STATUS:", response.status_code)
47
+ print("RESPONSE BODY:", response.json())
48
+
49
+ assert response.status_code == 200
50
+
51
+ # Vérifie la structure de la réponse
52
+ data = response.json()
53
+ assert "proba_bonne_estimation" in data
54
+ assert "prediction" in data
55
+ assert "seuil_applique" in data
56
+ assert data["fund_code"] == "FUND_001"
57
+
58
+ # La probabilité doit être entre 0 et 1
59
+ assert 0.0 <= data["proba_bonne_estimation"] <= 1.0
60
+
61
+
62
+ def test_predict_missing_required_field():
63
+ """
64
+ Test 3 : Vérifie que l'API renvoie une erreur 422 si un champ obligatoire (ex: quantity) est manquant.
65
+ """
66
+ invalid_payload = VALID_PAYLOAD.copy()
67
+ del invalid_payload["quantity"] # On supprime un champ obligatoire
68
+
69
+ response = client.post("/predict", json=invalid_payload)
70
+
71
+ # 422 Unprocessable Entity est le code standard de FastAPI pour une erreur de validation
72
+ assert response.status_code == 422
73
+
74
+
75
+ def test_predict_invalid_data_type():
76
+ """
77
+ Test 4 : Vérifie que l'API renvoie une erreur 422 si un type de donnée est incorrect
78
+ (ex: une chaîne de caractères au lieu d'un float pour 'quote').
79
+ """
80
+ invalid_payload = VALID_PAYLOAD.copy()
81
+ invalid_payload["quote"] = "Ceci_n_est_pas_un_chiffre"
82
+
83
+ response = client.post("/predict", json=invalid_payload)
84
+
85
+ assert response.status_code == 422