#!/usr/bin/env python3 """Train the TCN on real labeled windows (no synthetic data). Uses the windows extracted by extract_windows.py. Each window is a (200, 17) tensor labeled 0 (normal) or 1 (crash). The TCN learns to classify windows — detecting the real pre-crash microstructure pattern. Usage: # Train on crash day windows python scripts/train_tcn_windows.py --data data/windows/BTCUSDT_2021-05-19_windows.npz --out models/ --epochs 50 # Train on both crash + normal days (recommended) python scripts/train_tcn_windows.py \ --data data/windows/BTCUSDT_2021-05-19_windows.npz \ --data data/windows/BTCUSDT_2024-01-15_windows.npz \ --out models/ --epochs 50 --device cuda """ import argparse import logging import sys import time from pathlib import Path import numpy as np import torch import torch.nn as nn from torch.utils.data import DataLoader, TensorDataset ML_DIR = Path(__file__).resolve().parent.parent / "ml" sys.path.insert(0, str(ML_DIR)) PROJECT_ROOT = Path(__file__).resolve().parent.parent sys.path.insert(0, str(PROJECT_ROOT)) from flash_crash_watchdog.models.stage3_tcn import TCNDetector, TCNConfig logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") logger = logging.getLogger(__name__) class FocalLoss(nn.Module): """Focal loss for class imbalance — focuses training on hard examples.""" def __init__(self, alpha: float = 0.25, gamma: float = 2.0) -> None: super().__init__() self.alpha = alpha self.gamma = gamma def forward(self, preds: torch.Tensor, targets: torch.Tensor) -> torch.Tensor: # preds: (B, T) — per-timestep score in [0, 1] # targets: (B,) — window-level label # Use the score at the LAST timestep as the window-level prediction pred = preds[:, -1].squeeze() # (B,) BCE = nn.functional.binary_cross_entropy(pred, targets.float(), reduction="none") p_t = pred * targets + (1 - pred) * (1 - targets) alpha_t = self.alpha * targets + (1 - self.alpha) * (1 - targets) loss = alpha_t * (1 - p_t) ** self.gamma * BCE return loss.mean() def load_windows(paths: list[str]) -> tuple[np.ndarray, np.ndarray]: """Load and concatenate window data from multiple .npz files.""" all_windows = [] all_labels = [] for path in paths: data = np.load(path, allow_pickle=True) windows = data["windows"] labels = data["labels"] logger.info("Loaded %s: %d windows (%d positive)", path, len(windows), np.sum(labels)) all_windows.append(windows) all_labels.append(labels) windows = np.concatenate(all_windows, axis=0) labels = np.concatenate(all_labels, axis=0) logger.info("Total: %d windows (%d positive = %.1f%%)", len(windows), np.sum(labels), np.sum(labels) / len(labels) * 100) return windows, labels def train_tcn( windows: np.ndarray, labels: np.ndarray, epochs: int = 50, batch_size: int = 128, learning_rate: float = 1e-3, channels: int = 256, device: str = "cpu", ) -> TCNDetector: """Train the TCN on labeled windows. Args: windows: shape (N, window_size, 17) labels: shape (N,) — 0 or 1 """ logger.info("=" * 60) logger.info("TRAINING TCN ON REAL LABELED WINDOWS") logger.info(" Device: %s", device) logger.info(" Windows: %d", len(windows)) logger.info(" Positive: %d (%.1f%%)", np.sum(labels), np.sum(labels) / len(labels) * 100) logger.info(" Shape: %s", windows.shape) logger.info(" Epochs: %d", epochs) logger.info(" Batch: %d", batch_size) logger.info(" Channels: %d/layer", channels) logger.info("=" * 60) # Config window_size = windows.shape[1] input_dim = windows.shape[2] config = TCNConfig( num_channels=(channels,) * 8, kernel_size=3, input_dim=input_dim, dropout=0.1, sequence_length=window_size, ) # Model model = TCNDetector(config).to(device) optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=1e-4) criterion = FocalLoss(alpha=0.25, gamma=2.0) # Data — transpose to (N, input_dim, window_size) for Conv1d x = torch.FloatTensor(windows).permute(0, 2, 1).to(device) # (N, C, T) y = torch.LongTensor(labels).to(device) dataset = TensorDataset(x, y) # Split 80/20 n_train = int(len(dataset) * 0.8) n_val = len(dataset) - n_train train_ds, val_ds = torch.utils.data.random_split(dataset, [n_train, n_val]) train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, drop_last=True) val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False) logger.info("Train: %d windows | Val: %d windows", n_train, n_val) # Training loop best_val_loss = float("inf") history = {"train_loss": [], "val_loss": [], "val_acc": []} for epoch in range(epochs): t0 = time.time() model.train() train_loss = 0.0 n_batches = 0 for batch_x, batch_y in train_loader: optimizer.zero_grad() preds = model(batch_x) # (B, T) loss = criterion(preds, batch_y) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() train_loss += loss.item() n_batches += 1 train_loss /= max(1, n_batches) # Validation model.eval() val_loss = 0.0 val_correct = 0 val_total = 0 with torch.no_grad(): for batch_x, batch_y in val_loader: preds = model(batch_x) loss = criterion(preds, batch_y) val_loss += loss.item() # Predict: use last timestep score pred_label = (preds[:, -1].squeeze() > 0.5).long() val_correct += (pred_label == batch_y).sum().item() val_total += len(batch_y) val_loss /= max(1, len(val_loader)) val_acc = val_correct / max(1, val_total) history["train_loss"].append(train_loss) history["val_loss"].append(val_loss) history["val_acc"].append(val_acc) elapsed = time.time() - t0 if epoch % 5 == 0 or epoch == epochs - 1: logger.info("Epoch %3d/%d | train_loss=%.4f | val_loss=%.4f | val_acc=%.2f%% | %.1fs", epoch, epochs, train_loss, val_loss, val_acc * 100, elapsed) if val_loss < best_val_loss: best_val_loss = val_loss best_state = {k: v.clone() for k, v in model.state_dict().items()} # Restore best model.load_state_dict(best_state) logger.info("Best val_loss=%.4f, val_acc=%.2f%%", best_val_loss, max(history["val_acc"]) * 100) return model def main() -> int: parser = argparse.ArgumentParser(description="Train TCN on real labeled windows") parser.add_argument("--data", action="append", required=True, help="Path to .npz window file (can specify multiple)") parser.add_argument("--out", default="models/", help="Output directory") parser.add_argument("--epochs", type=int, default=50) parser.add_argument("--batch-size", type=int, default=128) parser.add_argument("--lr", type=float, default=1e-3) parser.add_argument("--channels", type=int, default=256, help="Channels per TCN layer (256 for A100, 64 for CPU)") parser.add_argument("--device", default="auto", help="cuda, cpu, or auto") args = parser.parse_args() # Device if args.device == "auto": device = "cuda" if torch.cuda.is_available() else "cpu" else: device = args.device if device == "cuda": logger.info("GPU: %s (%.1f GB)", torch.cuda.get_device_name(0), torch.cuda.get_device_properties(0).total_memory / 1e9) # Load windows windows, labels = load_windows(args.data) # Train model = train_tcn( windows, labels, epochs=args.epochs, batch_size=args.batch_size, learning_rate=args.lr, channels=args.channels, device=device, ) # Save out_dir = Path(args.out) out_dir.mkdir(parents=True, exist_ok=True) model_path = out_dir / "stage3_tcn_trained.pt" torch.save({ "model_state": model.state_dict(), "config": model.config, }, model_path) logger.info("Model saved to %s", model_path) return 0 if __name__ == "__main__": raise SystemExit(main())