Clasificador tabular de abandono de suscripción · 120K registros · MLP 3 capas · DDP-ready
prepare_data() descarga CSV, imputa nulos, normaliza con StandardScaler. setup() crea datasets 70/15/15 con clase pos ponderada.
training_step() calcula BCEWithLogitsLoss con pesos, validation_step() acumula F1/AUC/precision/recall.
ModelCheckpoint (monitor=val_f1_macro, mode=max) · EarlyStopping (patience=10) · LearningRateMonitor
Run nutripaw-churn-v1 · log automático de loss, F1, AUC, LR, GPU memory y confusion matrix al final.
# churn_model.py — NutriPaw Churn Classifier # Generado con skill: pytorch-lightning-entrenamiento-escalable import lightning as L import torch import torch.nn as nn import torch.nn.functional as F from torchmetrics.classification import BinaryF1Score, BinaryAUROC class ChurnMLP(nn.Module): """MLP tabular 3 capas: 12 → 256 → 128 → 64 → 2""" def __init__(self, input_dim: int = 12, dropout: float = 0.3): super().__init__() self.net = nn.Sequential( nn.Linear(input_dim, 256), nn.BatchNorm1d(256), nn.ReLU(), nn.Dropout(dropout), nn.Linear(256, 128), nn.BatchNorm1d(128), nn.ReLU(), nn.Dropout(dropout), nn.Linear(128, 64), nn.BatchNorm1d(64), nn.ReLU(), nn.Dropout(dropout), nn.Linear(64, 2), ) def forward(self, x): return self.net(x) class ChurnClassifier(L.LightningModule): """LightningModule para clasificación de churn en NutriPaw.""" def __init__( self, input_dim: int = 12, dropout: float = 0.3, lr: float = 3e-4, weight_decay: float = 1e-4, pos_weight: float = 4.56, # ratio negativo/positivo: 82%/18% ): super().__init__() self.save_hyperparameters() # serializa todos los args automáticamente self.model = ChurnMLP(input_dim, dropout) self.pos_weight = torch.tensor([pos_weight]) # TorchMetrics — reducción automática en DDP self.f1 = BinaryF1Score() self.auc = BinaryAUROC() # ── Training ─────────────────────────────────────────────────── def training_step(self, batch, batch_idx): x, y = batch logits = self.model(x) loss = F.cross_entropy(logits, y, weight=self.pos_weight.to(self.device)) self.log("train_loss", loss, prog_bar=True, on_step=False, on_epoch=True) return loss # ── Validation ───────────────────────────────────────────────── def validation_step(self, batch, batch_idx): x, y = batch logits = self.model(x) probs = torch.softmax(logits, dim=-1)[:, 1] preds = logits.argmax(dim=-1) val_loss = F.cross_entropy(logits, y) self.log_dict({ "val_loss": val_loss, "val_f1_macro": self.f1(preds, y), "val_auc": self.auc(probs, y), }, prog_bar=True, on_epoch=True) # ── Optimizer + Scheduler ────────────────────────────────────── def configure_optimizers(self): opt = torch.optim.AdamW( self.parameters(), lr=self.hparams.lr, weight_decay=self.hparams.weight_decay, ) sched = torch.optim.lr_scheduler.CosineAnnealingLR( opt, T_max=100, eta_min=1e-6 ) return {"optimizer": opt, "lr_scheduler": {"scheduler": sched, "interval": "epoch"}}
| Parámetro | Local (1 GPU) |
|---|---|
| accelerator | "gpu" |
| devices | 1 |
| strategy | "auto" |
| precision | "bf16-mixed" |
| max_epochs | 100 |
| gradient_clip_val | 1.0 |
| accumulate_grad_batches | 2 |
| log_every_n_steps | 50 |
| deterministic | True |
| fast_dev_run | False |
| Parámetro | Producción |
|---|---|
| accelerator | "gpu" |
| devices | 2 |
| strategy | "ddp" |
| precision | "bf16-mixed" |
| max_epochs | 100 |
| gradient_clip_val | 1.0 |
| accumulate_grad_batches | 1 |
| sync_batchnorm | True |
| find_unused_parameters | False |
| num_nodes | 1 |
monitor=val_f1_macro · mode=max · save_top_k=3 · dirpath=checkpoints/nutripaw/
monitor=val_loss · patience=10 · min_delta=1e-4 · mode=min
logging_interval=epoch · log_momentum=True
project=nutripaw-churn · name=nutripaw-churn-v1 · log_model=True
import lightning as L from lightning.pytorch.callbacks import ModelCheckpoint, EarlyStopping, LearningRateMonitor from lightning.pytorch.loggers import WandbLogger from churn_model import ChurnClassifier from churn_datamodule import ChurnDataModule L.seed_everything(42, workers=True) # reproducibilidad completa # ── Componentes ──────────────────────────────────────────────── model = ChurnClassifier(lr=3e-4, weight_decay=1e-4, pos_weight=4.56) dm = ChurnDataModule(csv_path="data/nutripaw_users.csv", batch_size=512) # ── Callbacks ───────────────────────────────────────────────── callbacks = [ ModelCheckpoint( dirpath="checkpoints/nutripaw", monitor="val_f1_macro", mode="max", save_top_k=3, filename="churn-{epoch:02d}-{val_f1_macro:.4f}", ), EarlyStopping(monitor="val_loss", patience=10, min_delta=1e-4), LearningRateMonitor(logging_interval="epoch"), ] # ── Logger ──────────────────────────────────────────────────── wandb_logger = WandbLogger(project="nutripaw-churn", name="nutripaw-churn-v1", log_model=True) # ── Trainer ─────────────────────────────────────────────────── trainer = L.Trainer( max_epochs=100, accelerator="gpu", devices=2, strategy="ddp", precision="bf16-mixed", gradient_clip_val=1.0, sync_batchnorm=True, # necesario en DDP con BatchNorm callbacks=callbacks, logger=wandb_logger, log_every_n_steps=50, deterministic=True, ) # ── Fit + Test ──────────────────────────────────────────────── trainer.fit(model, datamodule=dm) trainer.test(model, datamodule=dm, ckpt_path="best")
# Cargar mejor checkpoint y predecir model = ChurnClassifier.load_from_checkpoint( "checkpoints/nutripaw/churn-38-0.8470.ckpt" ) model.eval() trainer = L.Trainer(accelerator="gpu", devices=1) predictions = trainer.predict(model, datamodule=dm) # predictions: lista de tensores con probs probs = torch.cat(predictions) # [N, 2] churn_prob = probs[:, 1] # probabilidad de churn # Umbral de decision threshold = 0.40 # ajustado para max F1 at_risk = (churn_prob > threshold).sum() print(f"Usuarios en riesgo: {at_risk} / {len(churn_prob)}") # → Usuarios en riesgo: 6.302 / 35.000 (18.0%)