| |
| """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: |
| |
| |
| |
| pred = preds[:, -1].squeeze() |
| 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) |
|
|
| |
| 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 = 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) |
|
|
| |
| x = torch.FloatTensor(windows).permute(0, 2, 1).to(device) |
| y = torch.LongTensor(labels).to(device) |
| dataset = TensorDataset(x, y) |
|
|
| |
| 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) |
|
|
| |
| 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) |
| 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) |
|
|
| |
| 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() |
| |
| 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()} |
|
|
| |
| 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() |
|
|
| |
| 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) |
| |
|
|
| |
| windows, labels = load_windows(args.data) |
|
|
| |
| model = train_tcn( |
| windows, labels, |
| epochs=args.epochs, |
| batch_size=args.batch_size, |
| learning_rate=args.lr, |
| channels=args.channels, |
| device=device, |
| ) |
|
|
| |
| 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()) |
|
|