File size: 2,884 Bytes
52c1c67
 
850344d
52c1c67
850344d
3e9e69f
52c1c67
850344d
 
 
 
 
52c1c67
 
 
 
850344d
58debbc
3e9e69f
58debbc
3e9e69f
52c1c67
3e9e69f
58debbc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
52c1c67
 
850344d
52c1c67
 
850344d
58debbc
 
 
 
 
 
850344d
52c1c67
 
850344d
58debbc
 
 
 
 
 
 
 
 
 
 
 
 
850344d
 
 
52c1c67
850344d
52c1c67
850344d
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
from __future__ import annotations
"""
Predictive Model ๋ ˆ์ง€์ŠคํŠธ๋ฆฌ (ML/DL ๊ตฌํ˜„ ๊ต์ฒด ์ง€์ ).

predictive_model = ์•„๋ž˜ ์ธํ„ฐํŽ˜์ด์Šค๋ฅผ ๋งŒ์กฑํ•˜๋Š” ๊ฐ์ฒด/๋ชจ๋“ˆ (์‹œ๋‚˜๋ฆฌ์˜ค ๋ฌด๊ด€):
    predict(intent_id, features, *, training_data, dataset_path, model_prefix, train_params=None) -> float

- "predictive_model"์€ ML(sklearn)ยทDL(torch ๋“ฑ)์„ ์•„์šฐ๋ฅด๋Š” [2b] ์˜ˆ์ธก ๋ชจ๋ธ ๊ตฌํ˜„์„ ๊ฐ€๋ฆฌํ‚จ๋‹ค.
- config L2.model.predictive_model ๋กœ ์„ ํƒ (์—†์œผ๋ฉด "sklearn").
- ์ƒˆ ๊ตฌํ˜„์€ register_predictive_model("torch", <obj>)๋กœ ๋“ฑ๋ก.
- common.model_predict / GenericEngine.model_predict ๋Š” ํŠน์ • ๊ตฌํ˜„์„ ์ง์ ‘ importํ•˜์ง€ ์•Š๊ณ 
  ์ด ๋ ˆ์ง€์ŠคํŠธ๋ฆฌ๋กœ predictive_model์„ ๋ฐ›์•„ ํ˜ธ์ถœํ•œ๋‹ค.
"""
from typing import Any, Protocol


class PredictiveModel(Protocol):
    """์˜ˆ์ธก ๋ชจ๋ธ ๊ตฌํ˜„ ์ธํ„ฐํŽ˜์ด์Šค (sklearn/torchโ€ฆ ๊ณตํ†ต).

    intent๋ณ„ 0~1 ์ ์ˆ˜๋ฅผ ๋ฐ˜ํ™˜ํ•˜๋Š” predict ๋ฉ”์„œ๋“œ๋ฅผ ์ •์˜ํ•œ๋‹ค.
    """
    def predict(self, intent_id: str, features: dict[str, Any], *,
                training_data: dict, dataset_path, model_prefix: str,
                train_params: dict | None = None) -> float:
        """Intent์— ๋Œ€ํ•œ 0~1 ์˜ˆ์ธก ์ ์ˆ˜๋ฅผ ๋ฐ˜ํ™˜ํ•œ๋‹ค.

        Args:
            intent_id: ์˜ˆ์ธกํ•  Intent ID.
            features: ์ถ”๋ก ์— ์‚ฌ์šฉํ•  feature dict.
            training_data: ํ•™์Šต ๋ฐ์ดํ„ฐ(์‹œ๋‚˜๋ฆฌ์˜ค ์—”์ง„ ์ œ๊ณต).
            dataset_path: ์‹œ๋“œ ๋ฐ์ดํ„ฐ์…‹ ๊ฒฝ๋กœ.
            model_prefix: ์‹œ๋‚˜๋ฆฌ์˜ค๋ณ„ ๋ชจ๋ธ๋ช… ๋„ค์ž„์ŠคํŽ˜์ด์Šค.
            train_params: ํ•™์Šต ํ•˜์ดํผํŒŒ๋ผ๋ฏธํ„ฐ(config L2.model.train).
                ๋ฏธ์‚ฌ์šฉ ๊ตฌํ˜„์€ ๋ฌด์‹œ ๊ฐ€๋Šฅ.

        Returns:
            0~1 ๋ฒ”์œ„์˜ ์˜ˆ์ธก ์ ์ˆ˜.
        """
        ...


_PREDICTIVE_MODELS: dict[str, Any] = {}


def register_predictive_model(name: str, predictive_model: Any) -> None:
    """์˜ˆ์ธก ๋ชจ๋ธ ๊ตฌํ˜„์„ name์œผ๋กœ ๋“ฑ๋กํ•œ๋‹ค.

    Args:
        name: ๋“ฑ๋ก ํ‚ค (์˜ˆ: "torch").
        predictive_model: PredictiveModel ์ธํ„ฐํŽ˜์ด์Šค๋ฅผ ๋งŒ์กฑํ•˜๋Š” ๊ฐ์ฒด/๋ชจ๋“ˆ.
    """
    _PREDICTIVE_MODELS[name] = predictive_model


def get_predictive_model(name: str = "sklearn") -> Any:
    """name์— ๋“ฑ๋ก๋œ ์˜ˆ์ธก ๋ชจ๋ธ ๊ตฌํ˜„์„ ๋ฐ˜ํ™˜ํ•œ๋‹ค.

    "sklearn"์€ ์ตœ์ดˆ ํ˜ธ์ถœ ์‹œ ๊ธฐ๋ณธ ๊ตฌํ˜„์„ lazy ๋“ฑ๋กํ•œ๋‹ค.

    Args:
        name: ์กฐํšŒํ•  ์˜ˆ์ธก ๋ชจ๋ธ ๋“ฑ๋ก ํ‚ค.

    Returns:
        ๋“ฑ๋ก๋œ ์˜ˆ์ธก ๋ชจ๋ธ ๊ตฌํ˜„ ๊ฐ์ฒด/๋ชจ๋“ˆ.

    Raises:
        ValueError: ๋“ฑ๋ก๋˜์ง€ ์•Š์€ name์ด๋ฉด ๋ฐœ์ƒ.
    """
    if name == "sklearn" and "sklearn" not in _PREDICTIVE_MODELS:
        from models import sklearn_model            # ๊ธฐ๋ณธ ๊ตฌํ˜„ lazy ๋“ฑ๋ก
        _PREDICTIVE_MODELS["sklearn"] = sklearn_model
    try:
        return _PREDICTIVE_MODELS[name]
    except KeyError:
        raise ValueError(f"unknown predictive_model: {name!r} (๋“ฑ๋ก๋จ: {sorted(_PREDICTIVE_MODELS)})")