52 lines
1.5 KiB
Python
52 lines
1.5 KiB
Python
"""Map each experiment outcome to its initial model family."""
|
|
|
|
from typing import Literal
|
|
|
|
from pydantic import BaseModel, ConfigDict
|
|
|
|
from lmpm.training.targets import TargetTask, target_specs
|
|
|
|
ModelFamily = Literal["gaussian_process_regressor", "random_forest_classifier"]
|
|
|
|
|
|
class TargetModelRoute(BaseModel):
|
|
"""The default estimator family for one target and learning task."""
|
|
|
|
model_config = ConfigDict(frozen=True)
|
|
|
|
target_name: str
|
|
task: TargetTask
|
|
model_family: ModelFamily
|
|
|
|
|
|
def _default_route(target_name: str, task: TargetTask) -> TargetModelRoute:
|
|
if task == "regression":
|
|
return TargetModelRoute(
|
|
target_name=target_name,
|
|
task=task,
|
|
model_family="gaussian_process_regressor",
|
|
)
|
|
return TargetModelRoute(
|
|
target_name=target_name,
|
|
task=task,
|
|
model_family="random_forest_classifier",
|
|
)
|
|
|
|
|
|
DEFAULT_TARGET_MODEL_ROUTES = tuple(
|
|
_default_route(spec.name, spec.task) for spec in target_specs()
|
|
)
|
|
|
|
|
|
def model_route_for(target_name: str) -> TargetModelRoute:
|
|
"""Return the default compatible model route for a known outcome field."""
|
|
for route in DEFAULT_TARGET_MODEL_ROUTES:
|
|
if route.target_name == target_name:
|
|
return route
|
|
raise ValueError(f"unsupported training target: {target_name}")
|
|
|
|
|
|
def default_model_routes() -> tuple[TargetModelRoute, ...]:
|
|
"""Return every stable target-to-model mapping in dataset column order."""
|
|
return DEFAULT_TARGET_MODEL_ROUTES
|