Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 17 additions & 9 deletions const.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,23 @@
PROBA_THRESHOLD: float = 0.51
PROBA_MIN: float = 0.50


# MANAGEMENT COST

TRANSACTION_COST: float = 0.0010
MANAGEMENT_FEE_ANNUAL: float = 0.02
MIN_STOCKS_OPTIM: int = 3
MAX_STOCKS_SELECT: int = 10
WEIGHT_BOUNDS: tuple = (0.03, 0.20)

SHARPE_THRESHOLD: float = 0.30
MAX_DD_THRESHOLD: float = -0.40

BACKTEST_YEARS: int = 2


# FEATURES

FEATURE_COLS: list[str] = [
"rsi_lag1", "macd_lag1", "bb_low_lag1", "bb_mid_lag1",
"bb_high_lag1", "atr_lag1", "cluster_lag1",
Expand Down Expand Up @@ -64,15 +81,6 @@
"euro_volume", "volume", "open", "high", "low", "close",
]

TRANSACTION_COST: float = 0.0010
MIN_STOCKS_OPTIM: int = 3
MAX_STOCKS_SELECT: int = 10
WEIGHT_BOUNDS: tuple = (0.03, 0.20)

SHARPE_THRESHOLD: float = 0.30
MAX_DD_THRESHOLD: float = -0.40

BACKTEST_YEARS: int = 2

FEATURE_GROUPS = {
"momentum": [
Expand Down
68 changes: 34 additions & 34 deletions src/models/model_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,10 +19,8 @@
from pathlib import Path
from typing import Any, Optional

import mlflow
from dotenv import load_dotenv
from mlflow.exceptions import MlflowException
from mlflow.tracking import MlflowClient
from huggingface_hub import hf_hub_download

from const import MODEL_DIR
from src.utils.logger import setup_logger
Expand All @@ -32,50 +30,56 @@


# =============================================================================
# CONFIGURATION
# CONFIG
# =============================================================================

MLFLOW_TRACKING_URI = "https://soradata-alphaedge-registry.hf.space"
MLFLOW_USERNAME = "SORADATA"
CHAMPION_ALIAS = "champion"
HF_REPO_ID = "soradata/alphaedge-data"
LOCAL_MODEL_FILENAME = "ensemble_model.pkl"


HF_TOKEN = os.getenv("HF_TOKEN")
USE_MLFLOW = bool(HF_TOKEN)
USE_HF_HUB = bool(HF_TOKEN)

if USE_MLFLOW:
os.environ["MLFLOW_TRACKING_USERNAME"] = MLFLOW_USERNAME
os.environ["MLFLOW_TRACKING_PASSWORD"] = HF_TOKEN
mlflow.set_tracking_uri(MLFLOW_TRACKING_URI)
logger.info(f"MLflow activé pour le chargement du champion — tracking URI : {MLFLOW_TRACKING_URI}")
if USE_HF_HUB:
logger.info(f"Hugging Face activé pour le chargement du champion depuis : {HF_REPO_ID}")
else:
logger.warning("HF_TOKEN absent — chargement en mode local uniquement.")


# =============================================================================
# CHARGEMENT DEPUIS MLFLOW
# CHARGEMENT DEPUIS HUGGING FACE HUB
# =============================================================================

def _load_champion_from_mlflow(market_name: str) -> Optional[Any]:

def _load_champion_from_hf_hub(market_name: str) -> Optional[Any]:
"""
Charge le modèle aliasé 'champion' depuis le MLflow Model Registry.
Charge le modèle 'champion' depuis le dataset persistant Hugging Face.
Retourne None en cas d'échec (le fallback local prend alors le relais).
"""
registered_model_name = f"AlphaEdge_Ensemble_{market_name}"
model_uri = f"models:/{registered_model_name}@{CHAMPION_ALIAS}"
hf_filename = f"models/{market_name}/champion.pkl"
try:
model = mlflow.pyfunc.load_model(model_uri)
logger.info(f"[{market_name}] Champion chargé depuis MLflow : {model_uri}")
local_path = hf_hub_download(
repo_id=HF_REPO_ID,
repo_type="dataset",
filename=hf_filename,
token=HF_TOKEN
)

with open(local_path, "rb") as f:
model = pickle.load(f)

logger.info(f"[{market_name}] Champion chargé depuis Hugging Face Hub : {hf_filename}")
return model
except MlflowException as exc:
logger.warning(f"[{market_name}] Impossible de charger le champion MLflow ({model_uri}) : {exc}")
except Exception as exc:
logger.warning(
f"[{market_name}] Impossible de charger le champion depuis HF Hub ({hf_filename}) : {exc}"
)
return None


# =============================================================================
# CHARGEMENT DEPUIS LE FALLBACK LOCAL
# =============================================================================


def _local_model_path(market_name: str) -> Path:
return MODEL_DIR / market_name / LOCAL_MODEL_FILENAME

Expand All @@ -97,9 +101,6 @@ def _load_champion_from_local(market_name: str) -> Optional[Any]:
logger.info(f"[{market_name}] Modèle chargé depuis le fallback local : {local_path}")
return model
except (pickle.UnpicklingError, EOFError, AttributeError, ModuleNotFoundError) as exc:
# Ces erreurs signalent typiquement un fichier corrompu ou une
# incompatibilité de version entre l'environnement d'entraînement
# et celui d'inférence (classe déplacée/renommée, version sklearn...).
logger.error(f"[{market_name}] Fichier pickle illisible ou incompatible ({local_path}) : {exc}")
return None

Expand All @@ -121,15 +122,15 @@ def load_champion(market_name: str) -> Any:
Lève une exception si aucun modèle n'est disponible : le pipeline
ne doit jamais tourner sans modèle.
"""
model = _load_champion_from_mlflow(market_name) if USE_MLFLOW else None
model = _load_champion_from_hf_hub(market_name) if USE_HF_HUB else None

if model is None:
model = _load_champion_from_local(market_name)

if model is None:
raise RuntimeError(
f"[{market_name}] Aucun modèle champion disponible "
"(ni MLflow, ni local). Impossible de générer les signaux."
"(ni SUR hf hub, ni local). Impossible de générer les signaux."
)

return model
Expand All @@ -139,13 +140,12 @@ def clear_champion_cache(market_name: Optional[str] = None) -> None:
"""
Vide le cache de load_champion.

Utile après une nouvelle promotion (le champion vient de changer sur
MLflow) ou dans les tests, pour forcer un rechargement.
Note : lru_cache ne permet pas d'invalider une seule clé nativement,
donc on vide tout le cache quel que soit `market_name` fourni.
Utile après un nouvel entraînement local (le modèle vient de changer)
ou dans les tests, pour forcer un rechargement.

"""
load_champion.cache_clear()
if market_name:
logger.info(f"[{market_name}] Cache du champion invalidé.")
else:
logger.info("Cache du champion invalidé pour tous les marchés.")
logger.info("Cache du champion invalidé pour tous les marchés.")
18 changes: 18 additions & 0 deletions src/models/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
from src.models.ensemble import AlphaEdgeEnsemble
from src.utils.logger import setup_logger
from src.utils.metrics import calculate_financial_metrics
from huggingface_hub import HfApi


load_dotenv()
Expand Down Expand Up @@ -307,9 +308,26 @@ def _log_and_promote_to_mlflow(
)

if promote:
# Promote Mlflow ui
client.set_registered_model_alias(registered_model_name, "champion", model_version.version)
result.promoted = True
logger.info(f"[{market_name}] PROMOTION v{model_version.version} — {reason}")
try:
api = HfApi()
# Save on HF
local_model_path = MODEL_DIR / market_name / "ensemble_model.pkl"

api.upload_file(
path_or_fileobj=str(local_model_path),
path_in_repo=f"models/{market_name}/champion.pkl",
repo_id="soradata/alphaedge-data",
repo_type="dataset",
token=HF_TOKEN
)
logger.info(f"[{market_name}] Modèle persistant sauvegardé sur HF Hub ( soradata/alphaedge-data) ")
except Exception as e:
logger.error(f"[{market_name}] Sync faillure on Hf Hub : {e}")

else:
logger.warning(f"[{market_name}] CHALLENGER REJETÉ — {reason}")

Expand Down
Loading
Loading