File size: 8,513 Bytes
0f7aaa7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fe38ede
0f7aaa7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fe38ede
0f7aaa7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fe38ede
0f7aaa7
 
 
fe38ede
0f7aaa7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fe38ede
0f7aaa7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fe38ede
0f7aaa7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
"""
train_demand_model.py
---------------------
Entrena modelos Prophet de forecasting de demanda para cada producto
en sales_data y guarda las predicciones en demand_forecasts (Supabase).

Uso:
    python -m scripts.train_demand_model
    python -m scripts.train_demand_model --min-points 20   # mínimo de datos por producto
    python -m scripts.train_demand_model --dry-run         # solo muestra qué haría

Requisitos:
    pip install -r requirements-ml.txt
"""

import os
import sys
import logging
import argparse
from datetime import datetime, timezone

from dotenv import load_dotenv
load_dotenv()

# Silenciar logs de cmdstanpy/prophet/numpy
import warnings
warnings.filterwarnings("ignore")
logging.getLogger("cmdstanpy").setLevel(logging.CRITICAL)
logging.getLogger("prophet").setLevel(logging.CRITICAL)
logging.getLogger("numexpr").setLevel(logging.CRITICAL)

logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
logging.getLogger("httpx").setLevel(logging.WARNING)
logger = logging.getLogger(__name__)


def load_sales_data(db) -> list[dict]:
    """Carga todos los registros de sales_data desde Supabase."""
    result = db.table("sales_data").select(
        "date, product_name, quantity"
    ).order("date", desc=False).execute()
    return result.data or []


def group_by_product(records: list[dict]) -> dict[str, list[dict]]:
    """Agrupa registros por producto."""
    groups: dict[str, list[dict]] = {}
    for r in records:
        product = (r.get("product_name") or "desconocido").strip()
        if not product:
            continue
        groups.setdefault(product, []).append({
            "date": r["date"],
            "quantity": r["quantity"],
        })
    return groups


def save_forecasts(db, product: str, forecasts: dict, metrics: dict, n_records: int, dry_run: bool):
    """Guarda las predicciones de un producto en demand_forecasts."""
    run_at = datetime.now(timezone.utc).isoformat()
    rows = []
    for horizon, fc in forecasts.items():
        rows.append({
            "forecast_run_at": run_at,
            "product": product,
            "target_date": fc["target_date"],
            "horizon": horizon,
            "predicted_qty": fc["predicted_qty"],
            "lower_bound": fc["lower_bound"],
            "upper_bound": fc["upper_bound"],
            "model_name": "sales_demand_forecast",
            "data_points": n_records,
            "mape": metrics.get("mape"),
        })

    if dry_run:
        logger.info(f"  [DRY RUN] guardaría {len(rows)} filas para '{product}'")
        return

    db.table("demand_forecasts").insert(rows).execute()
    logger.info(f"  OK {len(rows)} forecasts guardados en Supabase para '{product}'")


def get_demand_model_id(db) -> str | None:
    try:
        r = db.table("ml_models").select("id").eq("type", "demand_forecast").execute()
        return r.data[0]["id"] if r.data else None
    except Exception:
        return None


def main():
    parser = argparse.ArgumentParser(description="Entrenamiento de modelos de demanda por producto")
    parser.add_argument("--min-points", type=int, default=30,
                        help="Mínimo de fechas únicas por producto para entrenar (default: 30)")
    parser.add_argument("--dry-run", action="store_true",
                        help="Muestra qué haría sin guardar nada")
    args = parser.parse_args()

    print(f"\n{'='*60}")
    print("  ENTRENAMIENTO -- Modelo de Forecasting de Demanda")
    print(f"  Mínimo de puntos por producto: {args.min_points}")
    print(f"  Modo: {'DRY RUN' if args.dry_run else 'PRODUCCIÓN'}")
    print(f"{'='*60}\n")

    from database.supabase_client import get_supabase
    from models.demand_forecast import DemandForecastModel

    db = get_supabase()

    # 1. Cargar datos
    logger.info("Cargando datos de sales_data desde Supabase...")
    records = load_sales_data(db)

    if not records:
        logger.error("No hay datos en sales_data. Importa tu Excel primero:")
        logger.error("  python -m scripts.import_excel --file tu_archivo.xlsx")
        sys.exit(1)

    logger.info(f"Total registros cargados: {len(records):,}")

    # 2. Agrupar por producto
    groups = group_by_product(records)
    logger.info(f"Productos encontrados: {len(groups)}")

    # 3. Filtrar por mínimo de puntos
    trainable = {p: recs for p, recs in groups.items() if len(recs) >= args.min_points}
    skipped = {p: len(recs) for p, recs in groups.items() if len(recs) < args.min_points}

    print(f"Productos con suficientes datos (>={args.min_points}): {len(trainable)}")
    if skipped:
        print(f"Productos omitidos por pocos datos: {len(skipped)}")
        for p, n in list(skipped.items())[:5]:
            print(f"  -- {p}: {n} registros")
        if len(skipped) > 5:
            print(f"  ... y {len(skipped)-5} más")
    print()

    if not trainable:
        logger.error(f"Ningún producto tiene suficientes datos (mínimo {args.min_points}).")
        logger.error("Reduce --min-points o importa más datos.")
        sys.exit(1)

    # 4. Entrenar modelo por producto
    results_summary = []
    total_success = 0
    total_failed = 0

    for product, product_records in trainable.items():
        print(f"  Entrenando: {product[:50]} ({len(product_records)} registros)...")

        model = DemandForecastModel(product_name=product)
        metrics = model.train(product_records)

        if "error" in metrics:
            logger.warning(f"  X Error en '{product}': {metrics['error']}")
            total_failed += 1
            continue

        try:
            forecasts = model.predict()
            save_forecasts(db, product, forecasts, metrics, model.n_records, args.dry_run)

            mape_str = f"{metrics['mape']:.1f}%" if metrics.get("mape") else "N/A"
            print(f"    MAPE: {mape_str} | "
                  f"sem1={forecasts['week_1']['predicted_qty']:.1f} | "
                  f"mes1={forecasts['month_1']['predicted_qty']:.1f} | "
                  f"mes2={forecasts['month_2']['predicted_qty']:.1f}")

            results_summary.append({
                "product": product,
                "mape": metrics.get("mape"),
                "week_1": forecasts["week_1"]["predicted_qty"],
                "month_1": forecasts["month_1"]["predicted_qty"],
            })
            total_success += 1

        except Exception as e:
            logger.error(f"  X Predicción falló para '{product}': {e}")
            total_failed += 1

    # 5. Actualizar ml_models y ml_model_runs
    if not args.dry_run and total_success > 0:
        try:
            avg_mape = None
            mapes = [r["mape"] for r in results_summary if r.get("mape")]
            if mapes:
                avg_mape = round(sum(mapes) / len(mapes), 2)

            model_id = get_demand_model_id(db)
            if model_id:
                db.table("ml_models").update({
                    "metrics": {
                        "mape": avg_mape,
                        "products_trained": total_success,
                        "last_run": datetime.now(timezone.utc).isoformat(),
                    },
                    "last_evaluated": datetime.now(timezone.utc).isoformat(),
                }).eq("id", model_id).execute()

                db.table("ml_model_runs").insert({
                    "model_id": model_id,
                    "metrics": {
                        "avg_mape": avg_mape,
                        "products_trained": total_success,
                        "products_failed": total_failed,
                    },
                    "data_source": "excel_import",
                    "rows_processed": len(records),
                    "status": "success",
                    "notes": f"{total_success} productos entrenados",
                }).execute()
        except Exception as e:
            logger.warning(f"No se pudo actualizar ml_models: {e}")

    # 6. Resumen final
    print(f"\n{'='*60}")
    print(f"  RESUMEN")
    print(f"  Productos entrenados: {total_success}")
    print(f"  Productos con error:  {total_failed}")
    if results_summary:
        mapes = [r["mape"] for r in results_summary if r.get("mape")]
        if mapes:
            print(f"  MAPE promedio: {sum(mapes)/len(mapes):.1f}%")
    print()
    print("  Siguiente paso: pregunta al bot en Telegram")
    print('  "¿Cuánto venderemos de [producto] el próximo mes?"')
    print(f"{'='*60}\n")


if __name__ == "__main__":
    main()