flash-crash-watchdog / scripts /train_tcn_windows.py
Dev2506's picture
Add files using upload-large-folder tool
2bbc43c verified
Raw
History Blame Contribute Delete
8.7 kB
#!/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())