CoolFace
Apppublic

matalee/hpce-dev

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
0likes
__init__.py76 linesDownload Raw Back to models
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