CoolFace
Apppublic

Kacemath/data-mining-tp

sourceHugging Faceupdated 7mo agoView on Hugging Face
0likes
preprocess.py158 linesDownload Raw Back to src
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