Kacemath/data-mining-tp
0
1from __future__ import annotations2 3import math4import pickle5from datetime import datetime6from pathlib import Path7from typing import Any8 9import numpy as np10import pandas as pd11 12LANGUAGE_MAPPING = {"en": 1, "zh": 2, "ja": 3}13 14PREFIX_TO_FORM_KEY = {15 "genres": "genres",16 "production_companies": "production_companies",17 "Keywords": "keywords",18 "cast": "cast",19}20 21 22def load_model(model_path: str | Path) -> Any:23 with Path(model_path).open("rb") as file:24 return pickle.load(file)25 26 27def get_model_feature_names(model: Any) -> list[str]:28 if not hasattr(model, "feature_names_in_"):29 raise ValueError("Model does not expose feature_names_in_.")30 return list(model.feature_names_in_)31 32 33def count_words(text: str | None) -> int:34 if text is None:35 return 036 normalized = str(text).strip()37 if not normalized:38 return 039 return len(normalized.split())40 41 42def runtime_category_code(runtime: float) -> int:43 if runtime < 90:44 return 045 if runtime < 120:46 return 147 return 248 49 50def parse_release_date(value: str | None) -> datetime:51 if not value:52 return datetime(2010, 1, 1)53 try:54 return datetime.strptime(value, "%Y-%m-%d")55 except ValueError as exc:56 raise ValueError("release_date must be in YYYY-MM-DD format.") from exc57 58 59def parse_feature_options(feature_names: list[str]) -> dict[str, list[str]]:60 options: dict[str, set[str]] = {k: set() for k in PREFIX_TO_FORM_KEY}61 62 for name in feature_names:63 for prefix in options:64 key = f"{prefix}_"65 if name.startswith(key) and name != f"{prefix}_other":66 options[prefix].add(name[len(key) :])67 68 return {k: sorted(v) for k, v in options.items()}69 70 71def _to_float(value: Any, default: float = 0.0) -> float:72 try:73 if value is None:74 return default75 return float(value)76 except (TypeError, ValueError):77 return default78 79 80def _to_int(value: Any, default: int = 0) -> int:81 try:82 if value is None:83 return default84 return int(value)85 except (TypeError, ValueError):86 return default87 88 89def build_feature_row(form_data: dict[str, Any], feature_names: list[str]) -> pd.DataFrame:90 row = {name: 0.0 for name in feature_names}91 92 budget = max(_to_float(form_data.get("budget"), 0.0), 0.0)93 popularity = max(_to_float(form_data.get("popularity"), 0.0), 0.0)94 runtime = max(_to_float(form_data.get("runtime"), 0.0), 0.0)95 96 release_date = parse_release_date(form_data.get("release_date"))97 release_season = ((release_date.month % 12) + 3) // 398 99 title_text = str(form_data.get("title") or "")100 tagline_text = str(form_data.get("tagline") or "")101 overview_text = str(form_data.get("overview") or "")102 103 values = {104 "belongs_to_collection": _to_int(form_data.get("belongs_to_collection"), 0),105 "homepage": _to_int(form_data.get("homepage"), 0),106 "has_tagline": _to_int(form_data.get("has_tagline"), 1 if tagline_text.strip() else 0),107 "original_language": LANGUAGE_MAPPING.get(str(form_data.get("original_language") or "").lower(), 0),108 "runtime": runtime,109 "num_of_cast": _to_float(form_data.get("num_of_cast"), 0.0),110 "num_of_crew": _to_float(form_data.get("num_of_crew"), 0.0),111 "gender_cast_1": _to_float(form_data.get("gender_cast_1"), 0.0),112 "gender_cast_2": _to_float(form_data.get("gender_cast_2"), 0.0),113 "count_cast_other": _to_float(form_data.get("count_cast_other"), 0.0),114 "title_word_count": _to_float(form_data.get("title_word_count"), count_words(title_text)),115 "tag_word_count": _to_float(form_data.get("tag_word_count"), count_words(tagline_text)),116 "overview_word_count": _to_float(form_data.get("overview_word_count"), count_words(overview_text)),117 "release_year": release_date.year,118 "release_month": release_date.month,119 "release_season": release_season,120 "runtime_category": runtime_category_code(runtime),121 "budget_log": math.log1p(budget),122 "popularity_log": math.log1p(popularity),123 }124 125 for key, value in values.items():126 if key in row:127 row[key] = value128 129 for prefix, form_key in PREFIX_TO_FORM_KEY.items():130 selected = form_data.get(form_key) or []131 if not isinstance(selected, list):132 selected = [selected]133 134 known = 0135 for item in selected:136 col = f"{prefix}_{item}"137 if col in row:138 row[col] = 1.0139 known += 1140 141 num_col = f"num_of_{prefix}"142 if num_col in row:143 row[num_col] = float(len(selected))144 145 other_col = f"{prefix}_other"146 if other_col in row:147 row[other_col] = 1.0 if len(selected) > known else 0.0148 149 df = pd.DataFrame([[row[name] for name in feature_names]], columns=feature_names)150 return df.replace([np.inf, -np.inf], 0).fillna(0)151 152 153def predict_revenue(model: Any, form_data: dict[str, Any]) -> float:154 feature_names = get_model_feature_names(model)155 frame = build_feature_row(form_data, feature_names)156 pred = model.predict(frame)[0]157 return float(pred)158 