matalee/hpce-dev
0
1from __future__ import annotations2"""3Predictive Model 레지스트리 (ML/DL 구현 교체 지점).4 5predictive_model = 아래 인터페이스를 만족하는 객체/모듈 (시나리오 무관):6 predict(intent_id, features, *, training_data, dataset_path, model_prefix, train_params=None) -> float7 8- "predictive_model"은 ML(sklearn)·DL(torch 등)을 아우르는 [2b] 예측 모델 구현을 가리킨다.9- config L2.model.predictive_model 로 선택 (없으면 "sklearn").10- 새 구현은 register_predictive_model("torch", <obj>)로 등록.11- common.model_predict / GenericEngine.model_predict 는 특정 구현을 직접 import하지 않고12 이 레지스트리로 predictive_model을 받아 호출한다.13"""14from typing import Any, Protocol15 16 17class PredictiveModel(Protocol):18 """예측 모델 구현 인터페이스 (sklearn/torch… 공통).19 20 intent별 0~1 점수를 반환하는 predict 메서드를 정의한다.21 """22 def predict(self, intent_id: str, features: dict[str, Any], *,23 training_data: dict, dataset_path, model_prefix: str,24 train_params: dict | None = None) -> float:25 """Intent에 대한 0~1 예측 점수를 반환한다.26 27 Args:28 intent_id: 예측할 Intent ID.29 features: 추론에 사용할 feature dict.30 training_data: 학습 데이터(시나리오 엔진 제공).31 dataset_path: 시드 데이터셋 경로.32 model_prefix: 시나리오별 모델명 네임스페이스.33 train_params: 학습 하이퍼파라미터(config L2.model.train).34 미사용 구현은 무시 가능.35 36 Returns:37 0~1 범위의 예측 점수.38 """39 ...40 41 42_PREDICTIVE_MODELS: dict[str, Any] = {}43 44 45def register_predictive_model(name: str, predictive_model: Any) -> None:46 """예측 모델 구현을 name으로 등록한다.47 48 Args:49 name: 등록 키 (예: "torch").50 predictive_model: PredictiveModel 인터페이스를 만족하는 객체/모듈.51 """52 _PREDICTIVE_MODELS[name] = predictive_model53 54 55def get_predictive_model(name: str = "sklearn") -> Any:56 """name에 등록된 예측 모델 구현을 반환한다.57 58 "sklearn"은 최초 호출 시 기본 구현을 lazy 등록한다.59 60 Args:61 name: 조회할 예측 모델 등록 키.62 63 Returns:64 등록된 예측 모델 구현 객체/모듈.65 66 Raises:67 ValueError: 등록되지 않은 name이면 발생.68 """69 if name == "sklearn" and "sklearn" not in _PREDICTIVE_MODELS:70 from models import sklearn_model # 기본 구현 lazy 등록71 _PREDICTIVE_MODELS["sklearn"] = sklearn_model72 try:73 return _PREDICTIVE_MODELS[name]74 except KeyError:75 raise ValueError(f"unknown predictive_model: {name!r} (등록됨: {sorted(_PREDICTIVE_MODELS)})")76 