Modelos de difusão interpretáveis para retornos financeiros sintéticos
Resumo
Este notebook descreve a adaptação de Diffusion-TS para gerar sequências sintéticas de retornos de ETF. O modelo de remoção de ruído prevê a sequência limpa como uma soma de tendência polinomial, componente sazonal de Fourier e resíduos, enquanto uma função de perda combinada nos domínios do tempo e de Fourier incentiva a fidelidade pontual e espectral. Um amostrador DDIM reduz o número de etapas reversas usadas na geração. Um classificador separado, treinado com sequências com ruído, orienta a geração de amostras rotuladas como regimes de baixa ou alta volatilidade.
O notebook avalia amostras incondicionais e condicionadas ao regime em comparação com dados históricos usando medidas de distribuição, correlação, autocorrelação e previsão posterior. Ele relata o uso de menor intensidade de orientação e maior temperatura de amostragem para o regime minoritário de alta volatilidade, a fim de evitar a perda de diversidade. O modelo é adaptado para retornos padronizados sem limites, sem o recorte usado em imagens. Os resultados dependem dos rótulos de regime simples de HMM e de configurações de orientação ajustadas; o notebook também observa o custo computacional de um modelo multiativos e não comprova que as sequências sintéticas capturem todas as propriedades de mercado necessárias para o trading.
Ideias principais
- Diffusion-TS prevê a sequência limpa de retornos e a decompõe em componentes de tendência, sazonais e residuais.
- Uma função de perda no domínio de Fourier incentiva a preservação da estrutura espectral e de autocorrelação.
- A amostragem DDIM reduz as etapas do processo reverso usadas para gerar sequências.
- Um classificador treinado com dados com ruído orienta a geração para regimes de volatilidade.
- As configurações de orientação por regime afetam tanto a fidelidade ao alvo quanto a diversidade das amostras, sobretudo em regimes minoritários.
Tags
Texto completo
# 05_diffusion_ts.py
```py
# ---
# jupyter:
# jupytext:
# cell_metadata_filter: tags,-all
# formats: ipynb,py:percent
# text_representation:
# extension: .py
# format_name: percent
# format_version: '1.3'
# jupytext_version: 1.19.3
# kernelspec:
# display_name: Python 3 (ipykernel)
# language: python
# name: python3
# ---
# %% [markdown]
# # Diffusion-TS: Interpretable Diffusion with Conditional Generation
#
# **Chapter 5: Synthetic Data Generation**
# **Section Reference**: Section 5.6 (Diffusion models for financial time series)
#
# **Docker image**: `ml4t-gpu`
#
# > **GPU recommended**: This notebook trains models with PyTorch/CUDA. It will run on CPU
# > but training may be very slow. For GPU acceleration:
# > ```bash
# > docker compose run --rm ml4t-gpu python 05_synthetic_data/05_diffusion_ts.py
# > ```
#
#
# ## Purpose
#
# This notebook implements **Diffusion-TS** (Yuan & Qiao, ICLR 2024), a diffusion
# model that decomposes the denoising prediction into **trend** (polynomial regression)
# and **seasonal** (Fourier basis) components. This interpretable structure encourages
# the model to separate slow drift from periodic patterns, analogous to classical STL
# decomposition but learned end-to-end within the diffusion framework.
#
# Unlike vanilla diffusion models that predict noise $\varepsilon$, Diffusion-TS
# predicts $x_0$ directly. The Fourier-domain loss further regularizes spectral
# fidelity -- critical for preserving autocorrelation structure in financial returns.
#
# ## Learning Objectives
#
# By completing this notebook, you will:
# - Implement a diffusion model with **interpretable trend+seasonal decomposition**
# - Train with a combined **time-domain + Fourier-domain** loss
# - Use **DDIM** fast sampling to reduce generation from 500 to 50 reverse steps
# - Build a **regime classifier** on noised data and apply **classifier guidance**
# to generate regime-conditional synthetic returns
# - Evaluate unconditional and conditional generation quality
#
# ## Cross-References
#
# - **Upstream**: ETF Universe loader (`data`)
# - **Downstream**: Regime-conditioned synthetic data for stress testing (Ch 20)
# - **Book**: Section 5.6 discusses diffusion and conditional generation
# - **Related**: [`02_tailgan_tail_risk`](02_tailgan_tail_risk.ipynb) (GAN), [`03_sigcwgan_signatures`](03_sigcwgan_signatures.ipynb) (GAN)
#
# ---
#
# ## From GANs to Interpretable Diffusion
#
# GANs face training instability, mode collapse, and limited interpretability.
# Diffusion models solve the first two by replacing adversarial training with a
# simple regression objective. Diffusion-TS goes further by decomposing the
# denoising network's output into components with known semantics:
#
# $$\hat{x}_0 = \text{Trend}(z) + \text{Season}(z) + \text{Residual}(z)$$
#
# where $z$ is the latent representation at diffusion step $t$. The trend block
# uses polynomial regression, the seasonal block selects top-$k$ Fourier modes,
# and the residual captures everything else. This makes the model's behavior
# inspectable -- you can visualize what the model attributes to drift versus
# cyclical patterns.
#
# ## References
#
# - **Paper**: Yuan, X. & Qiao, Y. (2024). "Diffusion-TS: Interpretable Diffusion
# for General Time Series Generation." ICLR 2024.
# - **Code**: https://github.com/Y-debug-sys/Diffusion-TS
#
# ## Key Adaptation Decisions
#
# 1. **No clamp(-1, 1)**: The original code clamps predicted $x_0$ to [-1, 1] for
# image-like data. Financial returns are StandardScaler-normalized (unbounded),
# so we remove the clamp entirely.
# 2. **Predict $x_0$** (not noise): The model outputs $\hat{x}_0 = \text{trend} +
# \text{season}$, then derives noise analytically. This pairs naturally with the
# decomposition.
# 3. **Fourier loss preserved**: The frequency-domain regularizer matches spectral
# properties -- essential for autocorrelation fidelity in returns.
# 4. **Classifier guidance**: A separate Transformer classifier trained on noised
# sequences enables regime-conditional generation at sampling time.
# %%
"""Diffusion-TS: Interpretable Diffusion with Conditional Generation."""
import json
import math
import warnings
from copy import deepcopy
from datetime import UTC, datetime
from pathlib import Path
# Scoped by category and module so a warning raised by this notebook's own code
# still reaches the reader.
warnings.filterwarnings("ignore", category=FutureWarning, module="torch")
warnings.filterwarnings("ignore", category=UserWarning, module="sklearn")
import matplotlib.pyplot as plt
import numpy as np
import polars as pl
import seaborn as sns
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange, reduce, repeat
from hmmlearn.hmm import GaussianHMM
from IPython.display import Image, display
from scipy import stats
from scipy.stats import kurtosis as calc_kurtosis
from sklearn.preprocessing import StandardScaler
from torch.utils.data import DataLoader, TensorDataset
from data import load_etfs
from utils.paths import get_chapter_dir, get_output_dir
from utils.reproducibility import set_global_seeds
from utils.style import COLORS, plot_fidelity_comparison, show_with_alt
# %% [markdown]
# ## Diffusion Process Overview
#
# Diffusion models work by gradually adding noise (forward process) and then
# learning to reverse it (denoising). Diffusion-TS adds interpretable
# trend+seasonal decomposition to the denoiser.
# %%
ASSETS_DIR = get_chapter_dir(5) / "assets"
if (ASSETS_DIR / "diffusion_forward_reverse.jpeg").exists():
display(Image(ASSETS_DIR / "diffusion_forward_reverse.jpeg", width=800))
# %% [markdown]
# ## Configuration
#
# The original Diffusion-TS paper uses seq_length=24, 6 stocks, and 10,000 epochs.
# We adapt for financial applications with the following considerations:
#
# | Parameter | Paper | Ours | Rationale |
# |-----------|-------|------|-----------|
# | seq_length | 24 | 60 | ~3 months of trading context |
# | feature_size | 6 | 20 | Balance diversity vs complexity |
# | epochs | 10,000 | 10,000 | Match paper for convergence |
# | lr | 1e-5→8e-4 | 1e-5→8e-4 | Warmup schedule from paper |
#
# With 20 assets × 60 timesteps = 1,200 values per sample (vs paper's 144),
# we train longer to capture the richer cross-asset structure.
# %% tags=["parameters"]
# Production defaults (Yuan & Qiao, ICLR 2024)
SEQ_LENGTH = 60 # ~3 months of trading context (paper uses 24)
FEATURE_SIZE = 20 # Number of ETFs (paper uses 6)
N_LAYER_DEC = 4 # Decoder layers (paper uses 2; we use 4 for more capacity)
TIMESTEPS = 500 # Forward diffusion steps
SAMPLING_TIMESTEPS = 50 # DDIM fast sampling steps
EPOCHS = 10000 # Training epochs (paper uses 10000)
BATCH_SIZE = 64 # Training batch size
WARMUP_STEPS = 500 # LR warmup steps
GRADIENT_ACCUMULATE_EVERY = 2 # Gradient accumulation steps
CLASSIFIER_EPOCHS = 1500 # Regime classifier training epochs
SCHEDULER_PATIENCE = 500 # ReduceLROnPlateau patience
N_SYNTHETIC = 500 # Number of unconditional synthetic sequences
N_COND = 100 # Number of conditional samples per regime
SEED = 42
# %% [markdown]
# Guidance is set per regime rather than globally. The minority high-volatility regime
# needs *gentler* steering than the majority one, which is the opposite of the intuition
# that a rarer target needs a harder push: classifier gradients toward a rare class are
# strong enough to drive every sample to extreme volatility, collapsing the mode the
# guidance was meant to reach. Its lower scale is paired with a higher temperature, which
# buys back diversity, and a larger eta, which leaves more noise in each sampling step.
# %%
set_global_seeds(SEED)
# Configuration
RETRAIN = False # Set True to retrain even if checkpoint exists
CONFIG = {
# Sequence/architecture - scaled from paper (24×6) to financial use case
"seq_length": SEQ_LENGTH,
"feature_size": FEATURE_SIZE,
"d_model": 64,
"n_heads": 4,
"n_layer_enc": 2,
"n_layer_dec": N_LAYER_DEC,
# Diffusion process
"timesteps": TIMESTEPS,
"sampling_timesteps": SAMPLING_TIMESTEPS,
"eta": 0.0, # Deterministic DDIM (eta=1 made variance worse, not better)
"beta_schedule": "cosine",
"loss_type": "l1",
# Training - match paper's schedule
"epochs": EPOCHS,
"batch_size": BATCH_SIZE,
"lr": 1e-5, # Paper's base_lr (not warmup target)
"warmup_lr": 8e-4, # Paper's warmup target
"warmup_steps": WARMUP_STEPS,
"ema_decay": 0.995,
"gradient_accumulate_every": GRADIENT_ACCUMULATE_EVERY,
# Data split
"start_date": "2005-01-01",
"holdout_start": "2024-01-01",
# Classifier guidance for regime-conditional generation
"classifier_epochs": CLASSIFIER_EPOCHS,
"classifier_lr": 5e-4,
# Per-regime guidance; see the markdown above for why the minority class is
# steered less hard than the majority one.
"guidance_settings": {
0: {"scale": 0.75, "temperature": 1.0, "eta": 0.5}, # Low-Vol
1: {"scale": 0.3, "temperature": 2.0, "eta": 0.7}, # High-Vol: gentler
},
"n_regimes": 2, # Low-Vol vs High-Vol (imbalanced data precludes 3-way)
}
# %%
# Output paths and reproducibility
OUTPUT_DIR = get_output_dir(5, "diffusion_ts")
CHECKPOINT_DIR = OUTPUT_DIR / "checkpoints"
CHECKPOINT_DIR.mkdir(parents=True, exist_ok=True)
CHECKPOINT_PATH = CHECKPOINT_DIR / "diffusion_ts_model.pt"
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")
# %% [markdown]
# ## 1. Data Loading and Preparation
#
# We load daily ETF returns and create overlapping sequences. The temporal
# train/holdout split at 2024-01-01 ensures unbiased TSTR evaluation.
# %%
def load_returns_data(start_date: str, n_assets: int) -> tuple[np.ndarray, np.ndarray, list[str]]:
"""Load daily returns for ETF assets with temporal split."""
df = load_etfs()
start_dt = pl.lit(start_date).str.to_date()
returns_df = (
df.filter(pl.col("timestamp") >= start_dt)
.sort(["symbol", "timestamp"])
.with_columns(pl.col("close").pct_change().over("symbol").alias("return"))
.pivot(on="symbol", index="timestamp", values="return")
.sort("timestamp")
.drop_nulls()
)
data_cols = [c for c in returns_df.columns if c != "timestamp"][:n_assets]
timestamps = returns_df.select("timestamp").to_numpy().flatten()
returns = returns_df.select(data_cols).to_numpy().astype(np.float32)
print(f"Loaded {len(returns)} days of returns for {len(data_cols)} assets")
print(f"Date range: {timestamps[0]} to {timestamps[-1]}")
return returns, timestamps, data_cols
# %% [markdown]
# ### Create Overlapping Sequences
#
# Sliding windows maximize sample count from limited financial data.
# %%
def create_sequences(data: np.ndarray, seq_length: int) -> np.ndarray:
"""Create overlapping sequences from time series data."""
n_seq = len(data) - seq_length + 1
seqs = np.zeros((n_seq, seq_length, data.shape[1]), dtype=np.float32)
for i in range(n_seq):
seqs[i] = data[i : i + seq_length]
return seqs
# %%
all_returns, all_timestamps, asset_names = load_returns_data(
CONFIG["start_date"], CONFIG["feature_size"]
)
n_assets = all_returns.shape[1]
# Temporal split
holdout_dt = np.datetime64(CONFIG["holdout_start"])
train_mask = all_timestamps < holdout_dt
returns = all_returns[train_mask]
holdout_returns = all_returns[~train_mask]
print(f"\nTrain: {len(returns):,} days | Holdout: {len(holdout_returns):,} days")
sequences = create_sequences(returns, CONFIG["seq_length"])
holdout_sequences = create_sequences(holdout_returns, CONFIG["seq_length"])
print(f"Train sequences: {sequences.shape} | Holdout: {holdout_sequences.shape}")
# %% [markdown]
# ## 2. Utility Modules
#
# These building blocks support the Transformer architecture: sinusoidal
# timestep embeddings, learnable positional encoding, and adaptive layer
# normalization that conditions on the diffusion timestep.
# %%
# Inline helpers (from Diffusion-TS model_utils)
def extract(a, t, x_shape):
"""Gather values from `a` at indices `t`, reshape for broadcasting."""
b, *_ = t.shape
out = a.gather(-1, t)
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
# %%
class SinusoidalPosEmb(nn.Module):
"""Sinusoidal positional embedding for diffusion timestep."""
def __init__(self, dim):
super().__init__()
self.dim = dim
def forward(self, x):
half_dim = self.dim // 2
emb = math.log(10000) / (half_dim - 1)
emb = torch.exp(torch.arange(half_dim, device=x.device) * -emb)
emb = x[:, None] * emb[None, :]
return torch.cat((emb.sin(), emb.cos()), dim=-1)
# %% [markdown]
# ### Adaptive Layer Normalization
#
# AdaLayerNorm modulates the normalized activations by a scale and shift
# derived from the diffusion timestep embedding. This gives each Transformer
# layer information about the current noise level.
# %%
class AdaLayerNorm(nn.Module):
"""Layer norm conditioned on diffusion timestep via scale+shift."""
def __init__(self, n_embd):
super().__init__()
self.emb = SinusoidalPosEmb(n_embd)
self.silu = nn.SiLU()
self.linear = nn.Linear(n_embd, n_embd * 2)
self.layernorm = nn.LayerNorm(n_embd, elementwise_affine=False)
def forward(self, x, timestep, label_emb=None):
emb = self.emb(timestep)
if label_emb is not None:
emb = emb + label_emb
emb = self.linear(self.silu(emb)).unsqueeze(1)
scale, shift = torch.chunk(emb, 2, dim=2)
return self.layernorm(x) * (1 + scale) + shift
# %% [markdown]
# ### Learnable Positional Encoding and Conv Embedding
#
# LearnablePositionalEncoding adds a learned position vector to each timestep.
# Conv_MLP projects the input feature dimension to the model dimension using
# a 1D convolution, which captures local patterns in the feature axis.
# %%
class LearnablePositionalEncoding(nn.Module):
"""Learned positional encoding added to sequence embeddings."""
def __init__(self, d_model, dropout=0.1, max_len=1024):
super().__init__()
self.dropout = nn.Dropout(p=dropout)
self.pe = nn.Parameter(torch.empty(1, max_len, d_model))
nn.init.uniform_(self.pe, -0.02, 0.02)
def forward(self, x):
return self.dropout(x + self.pe[:, : x.size(1)])
# %%
class Transpose(nn.Module):
"""Transpose wrapper for use in nn.Sequential."""
def __init__(self, shape: tuple):
super().__init__()
self.shape = shape
def forward(self, x):
return x.transpose(*self.shape)
# %%
class Conv_MLP(nn.Module):
"""1D convolution embedding: features → model dimension."""
def __init__(self, in_dim, out_dim, resid_pdrop=0.0):
super().__init__()
self.sequential = nn.Sequential(
Transpose(shape=(1, 2)),
nn.Conv1d(in_dim, out_dim, 3, stride=1, padding=1),
nn.Dropout(p=resid_pdrop),
)
def forward(self, x):
return self.sequential(x).transpose(1, 2)
# %% [markdown]
# ## 3. Interpretable Decomposition
#
# The key innovation of Diffusion-TS: each decoder layer extracts **trend**
# and **seasonal** components from its intermediate representation.
#
# - **TrendBlock**: Learns a polynomial basis (degree 3) via 1D convolutions,
# then multiplies by a polynomial space $[t, t^2, t^3]$ to produce a smooth trend.
# - **FourierLayer**: Computes the DFT of the latent, selects top-$k$ frequencies
# by magnitude, then reconstructs via inverse DFT. This extracts dominant
# periodic patterns without assuming a fixed period.
#
# Across decoder layers, trend and seasonal residuals accumulate, building up
# the full decomposition progressively.
# %%
class TrendBlock(nn.Module):
"""Polynomial regression on latent representation → smooth trend."""
def __init__(self, in_dim, out_dim, in_feat, out_feat, act):
super().__init__()
trend_poly = 3
self.trend = nn.Sequential(
nn.Conv1d(in_channels=in_dim, out_channels=trend_poly, kernel_size=3, padding=1),
act,
Transpose(shape=(1, 2)),
nn.Conv1d(in_feat, out_feat, 3, stride=1, padding=1),
)
lin_space = torch.arange(1, out_dim + 1, 1) / (out_dim + 1)
self.poly_space = torch.stack([lin_space ** float(p + 1) for p in range(trend_poly)], dim=0)
def forward(self, x):
x = self.trend(x).transpose(1, 2)
trend_vals = torch.matmul(x.transpose(1, 2), self.poly_space.to(x.device))
return trend_vals.transpose(1, 2)
# %% [markdown]
# ### FourierLayer: Top-k Frequency Selection
#
# Rather than using all Fourier coefficients, the layer selects the top-$k$
# frequencies (by magnitude) and reconstructs only those. This acts as a
# learned bandpass filter that adapts to the data's spectral content.
# %%
class FourierLayer(nn.Module):
"""Extract seasonal component via top-k inverse DFT."""
def __init__(self, d_model, low_freq=1, factor=1):
super().__init__()
self.d_model = d_model
self.factor = factor
self.low_freq = low_freq
def forward(self, x):
"""x: (b, t, d)"""
b, t, d = x.shape
x_freq = torch.fft.rfft(x, dim=1)
if t % 2 == 0:
x_freq = x_freq[:, self.low_freq : -1]
f = torch.fft.rfftfreq(t)[self.low_freq : -1]
else:
x_freq = x_freq[:, self.low_freq :]
f = torch.fft.rfftfreq(t)[self.low_freq :]
x_freq, index_tuple = self.topk_freq(x_freq)
f = repeat(f, "f -> b f d", b=x_freq.size(0), d=x_freq.size(2)).to(x_freq.device)
f = rearrange(f[index_tuple], "b f d -> b f () d").to(x_freq.device)
return self.extrapolate(x_freq, f, t)
def extrapolate(self, x_freq, f, t):
x_freq = torch.cat([x_freq, x_freq.conj()], dim=1)
f = torch.cat([f, -f], dim=1)
t_range = rearrange(torch.arange(t, dtype=torch.float), "t -> () () t ()").to(x_freq.device)
amp = rearrange(x_freq.abs(), "b f d -> b f () d")
phase = rearrange(x_freq.angle(), "b f d -> b f () d")
x_time = amp * torch.cos(2 * math.pi * f * t_range + phase)
return reduce(x_time, "b f t d -> b t d", "sum")
def topk_freq(self, x_freq):
length = x_freq.shape[1]
top_k = int(self.factor * math.log(max(length, 2)))
top_k = max(top_k, 1)
values, indices = torch.topk(x_freq.abs(), top_k, dim=1, largest=True, sorted=True)
mesh_a, mesh_b = torch.meshgrid(
torch.arange(x_freq.size(0)), torch.arange(x_freq.size(2)), indexing="ij"
)
index_tuple = (mesh_a.unsqueeze(1), indices, mesh_b.unsqueeze(1))
x_freq = x_freq[index_tuple]
return x_freq, index_tuple
# %% [markdown]
# ## 4. Transformer Architecture
#
# The encoder-decoder Transformer processes noised sequences. The **encoder**
# creates a contextual representation conditioned on the diffusion timestep.
# The **decoder** cross-attends to the encoder output and, at each layer,
# extracts trend and seasonal residuals that accumulate across layers.
# %%
class FullAttention(nn.Module):
"""Multi-head self-attention."""
def __init__(self, n_embd, n_head, attn_pdrop=0.0, resid_pdrop=0.0):
super().__init__()
assert n_embd % n_head == 0
self.key = nn.Linear(n_embd, n_embd)
self.query = nn.Linear(n_embd, n_embd)
self.value = nn.Linear(n_embd, n_embd)
self.attn_drop = nn.Dropout(attn_pdrop)
self.resid_drop = nn.Dropout(resid_pdrop)
self.proj = nn.Linear(n_embd, n_embd)
self.n_head = n_head
def forward(self, x, mask=None):
B, T, C = x.size()
k = self.key(x).view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
q = self.query(x).view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
v = self.value(x).view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(k.size(-1)))
att = F.softmax(att, dim=-1)
att = self.attn_drop(att)
y = att @ v
y = y.transpose(1, 2).contiguous().view(B, T, C)
return self.resid_drop(self.proj(y))
# %% [markdown]
# ### Cross-Attention
#
# The decoder cross-attends to the encoder's output. Queries come from the
# decoder, while keys and values come from the encoder -- this lets each
# decoder position gather relevant context from the full encoded sequence.
# %%
class CrossAttention(nn.Module):
"""Multi-head cross-attention (decoder queries, encoder keys/values)."""
def __init__(self, n_embd, condition_embd, n_head, attn_pdrop=0.0, resid_pdrop=0.0):
super().__init__()
assert n_embd % n_head == 0
self.key = nn.Linear(condition_embd, n_embd)
self.query = nn.Linear(n_embd, n_embd)
self.value = nn.Linear(condition_embd, n_embd)
self.attn_drop = nn.Dropout(attn_pdrop)
self.resid_drop = nn.Dropout(resid_pdrop)
self.proj = nn.Linear(n_embd, n_embd)
self.n_head = n_head
def forward(self, x, encoder_output, mask=None):
B, T, C = x.size()
B, T_E, _ = encoder_output.size()
k = self.key(encoder_output).view(B, T_E, self.n_head, C // self.n_head).transpose(1, 2)
q = self.query(x).view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
v = self.value(encoder_output).view(B, T_E, self.n_head, C // self.n_head).transpose(1, 2)
att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(k.size(-1)))
att = F.softmax(att, dim=-1)
att = self.attn_drop(att)
y = att @ v
y = y.transpose(1, 2).contiguous().view(B, T, C)
return self.resid_drop(self.proj(y))
# %% [markdown]
# ### Encoder and Decoder Blocks
#
# Each encoder block applies adaptive layer norm → self-attention → FFN.
# Each decoder block adds cross-attention to the encoder output, then
# splits the representation into two branches: one for trend extraction
# (polynomial regression) and one for seasonal extraction (Fourier layer).
# %%
class EncoderBlock(nn.Module):
"""Transformer encoder block with timestep-conditioned AdaLayerNorm."""
def __init__(self, n_embd=64, n_head=4, attn_pdrop=0.0, resid_pdrop=0.0, mlp_hidden_times=4):
super().__init__()
self.ln1 = AdaLayerNorm(n_embd)
self.ln2 = nn.LayerNorm(n_embd)
self.attn = FullAttention(n_embd, n_head, attn_pdrop, resid_pdrop)
self.mlp = nn.Sequential(
nn.Linear(n_embd, mlp_hidden_times * n_embd),
nn.GELU(),
nn.Linear(mlp_hidden_times * n_embd, n_embd),
nn.Dropout(resid_pdrop),
)
def forward(self, x, timestep, mask=None):
x = x + self.attn(self.ln1(x, timestep), mask=mask)
x = x + self.mlp(self.ln2(x))
return x
# %%
class Encoder(nn.Module):
"""Stack of encoder blocks."""
def __init__(self, n_layer=2, n_embd=64, n_head=4, attn_pdrop=0.0, resid_pdrop=0.0):
super().__init__()
self.blocks = nn.ModuleList(
[EncoderBlock(n_embd, n_head, attn_pdrop, resid_pdrop) for _ in range(n_layer)]
)
def forward(self, x, t):
for block in self.blocks:
x = block(x, t)
return x
# %%
class DecoderBlock(nn.Module):
"""Decoder block: self-attn + cross-attn + trend/season extraction."""
def __init__(
self,
n_channel,
n_feat,
n_embd=64,
n_head=4,
attn_pdrop=0.0,
resid_pdrop=0.0,
mlp_hidden_times=4,
condition_dim=64,
):
super().__init__()
self.ln1 = AdaLayerNorm(n_embd)
self.ln2 = nn.LayerNorm(n_embd)
self.ln1_1 = AdaLayerNorm(n_embd)
self.attn1 = FullAttention(n_embd, n_head, attn_pdrop, resid_pdrop)
self.attn2 = CrossAttention(n_embd, condition_dim, n_head, attn_pdrop, resid_pdrop)
act = nn.GELU()
self.trend = TrendBlock(n_channel, n_channel, n_embd, n_feat, act=act)
self.seasonal = FourierLayer(d_model=n_embd)
self.mlp = nn.Sequential(
nn.Linear(n_embd, mlp_hidden_times * n_embd),
nn.GELU(),
nn.Linear(mlp_hidden_times * n_embd, n_embd),
nn.Dropout(resid_pdrop),
)
self.proj = nn.Conv1d(n_channel, n_channel * 2, 1)
self.linear = nn.Linear(n_embd, n_feat)
def forward(self, x, encoder_output, timestep, mask=None):
x = x + self.attn1(self.ln1(x, timestep), mask=mask)
x = x + self.attn2(self.ln1_1(x, timestep), encoder_output, mask=mask)
x1, x2 = self.proj(x).chunk(2, dim=1)
trend, season = self.trend(x1), self.seasonal(x2)
x = x + self.mlp(self.ln2(x))
m = torch.mean(x, dim=1, keepdim=True)
return x - m, self.linear(m), trend, season
# %%
class Decoder(nn.Module):
"""Stack of decoder blocks, accumulating trend and seasonal components."""
def __init__(
self,
n_channel,
n_feat,
n_embd=64,
n_head=4,
n_layer=4,
attn_pdrop=0.0,
resid_pdrop=0.0,
condition_dim=64,
):
super().__init__()
self.d_model = n_embd
self.n_feat = n_feat
self.blocks = nn.ModuleList(
[
DecoderBlock(
n_feat=n_feat,
n_channel=n_channel,
n_embd=n_embd,
n_head=n_head,
attn_pdrop=attn_pdrop,
resid_pdrop=resid_pdrop,
condition_dim=condition_dim,
)
for _ in range(n_layer)
]
)
def forward(self, x, t, enc):
b, c, _ = x.shape
mean = []
season = torch.zeros((b, c, self.d_model), device=x.device)
trend = torch.zeros((b, c, self.n_feat), device=x.device)
for block in self.blocks:
x, residual_mean, residual_trend, residual_season = block(x, enc, t)
season += residual_season
trend += residual_trend
mean.append(residual_mean)
mean = torch.cat(mean, dim=1)
return x, mean, trend, season
# %% [markdown]
# ### Full Transformer
#
# The Transformer wraps encoder and decoder with input/output projections.
# The forward pass returns two tensors whose sum is $\hat{x}_0$: `trend`, which is the
# trend accumulated across decoder layers plus the residual's mean, and `season_error`,
# which is the seasonal component projected back to feature space plus the residual with
# that mean removed. Splitting the residual this way keeps the trend term carrying the
# level and the seasonal term carrying the variation around it.
# %%
class DiffusionTransformer(nn.Module):
"""Encoder-decoder Transformer with interpretable trend+seasonal decomposition."""
def __init__(
self,
n_feat,
n_channel,
n_layer_enc=2,
n_layer_dec=4,
n_embd=64,
n_heads=4,
attn_pdrop=0.0,
resid_pdrop=0.0,
mlp_hidden_times=4,
max_len=2048,
):
super().__init__()
self.emb = Conv_MLP(n_feat, n_embd, resid_pdrop=resid_pdrop)
self.inverse = Conv_MLP(n_embd, n_feat, resid_pdrop=resid_pdrop)
kernel_size, padding = (1, 0) if n_feat < 32 and n_channel < 64 else (5, 2)
self.combine_s = nn.Conv1d(
n_embd,
n_feat,
kernel_size=kernel_size,
stride=1,
padding=padding,
padding_mode="circular",
bias=False,
)
self.combine_m = nn.Conv1d(
n_layer_dec,
1,
kernel_size=1,
stride=1,
padding=0,
padding_mode="circular",
bias=False,
)
self.encoder = Encoder(n_layer_enc, n_embd, n_heads, attn_pdrop, resid_pdrop)
self.pos_enc = LearnablePositionalEncoding(n_embd, dropout=resid_pdrop, max_len=max_len)
self.decoder = Decoder(
n_channel,
n_feat,
n_embd,
n_heads,
n_layer_dec,
attn_pdrop,
resid_pdrop,
condition_dim=n_embd,
)
self.pos_dec = LearnablePositionalEncoding(n_embd, dropout=resid_pdrop, max_len=max_len)
def forward(self, x, t, return_res=False):
emb = self.emb(x)
inp_enc = self.pos_enc(emb)
enc_cond = self.encoder(inp_enc, t)
inp_dec = self.pos_dec(emb)
output, mean, trend, season = self.decoder(inp_dec, t, enc_cond)
res = self.inverse(output)
res_m = torch.mean(res, dim=1, keepdim=True)
season_error = self.combine_s(season.transpose(1, 2)).transpose(1, 2) + res - res_m
trend = self.combine_m(mean) + res_m + trend
if return_res:
return trend, self.combine_s(season.transpose(1, 2)).transpose(1, 2), res - res_m
return trend, season_error
# %% [markdown]
# ## 5. Diffusion Process
#
# The `DiffusionTS` class implements the full DDPM framework with $x_0$
# prediction. Key methods:
# - `q_sample`: Forward process -- add noise to clean data
# - `model_predictions`: Get model's $\hat{x}_0$ and derived noise
# - `p_mean_variance`: Compute posterior $p(x_{t-1}|x_t)$ -- **no clamp**
# - the training loss, internally: L1 in the time domain plus a Fourier term in the
# frequency domain
# - `fast_sample`: DDIM-style accelerated sampling
#
# **Critical adaptation**: We remove `clamp(-1, 1)` from `p_mean_variance`
# and `model_predictions`. Financial returns are StandardScaler-normalized
# but unbounded -- clamping would truncate the distribution tails.
# %%
def cosine_beta_schedule(timesteps, s=0.008):
"""Cosine schedule (Nichol & Dhariwal 2021)."""
steps = timesteps + 1
x = torch.linspace(0, timesteps, steps, dtype=torch.float64)
alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * math.pi * 0.5) ** 2
alphas_cumprod = alphas_cumprod / alphas_cumprod[0]
betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])
return torch.clip(betas, 0, 0.999)
# %%
class DiffusionTS(nn.Module):
"""Diffusion-TS: x_0-prediction diffusion with Fourier loss."""
def __init__(
self,
seq_length,
feature_size,
n_layer_enc=2,
n_layer_dec=4,
d_model=64,
timesteps=500,
sampling_timesteps=None,
loss_type="l1",
n_heads=4,
mlp_hidden_times=4,
eta=0.0,
reg_weight=None,
):
super().__init__()
self.eta = eta
self.seq_length = seq_length
self.feature_size = feature_size
self.ff_weight = reg_weight if reg_weight is not None else math.sqrt(seq_length) / 5
self.model = DiffusionTransformer(
n_feat=feature_size,
n_channel=seq_length,
n_layer_enc=n_layer_enc,
n_layer_dec=n_layer_dec,
n_heads=n_heads,
mlp_hidden_times=mlp_hidden_times,
max_len=seq_length,
n_embd=d_model,
)
betas = cosine_beta_schedule(timesteps)
alphas = 1.0 - betas
alphas_cumprod = torch.cumprod(alphas, dim=0)
alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value=1.0)
self.num_timesteps = int(timesteps)
self.loss_type = loss_type
self.sampling_timesteps = (
sampling_timesteps if sampling_timesteps is not None else timesteps
)
self.fast_sampling = self.sampling_timesteps < timesteps
def register(name, val):
self.register_buffer(name, val.to(torch.float32))
register("betas", betas)
register("alphas_cumprod", alphas_cumprod)
register("alphas_cumprod_prev", alphas_cumprod_prev)
register("sqrt_alphas_cumprod", torch.sqrt(alphas_cumprod))
register("sqrt_one_minus_alphas_cumprod", torch.sqrt(1.0 - alphas_cumprod))
register("log_one_minus_alphas_cumprod", torch.log(1.0 - alphas_cumprod))
register("sqrt_recip_alphas_cumprod", torch.sqrt(1.0 / alphas_cumprod))
register("sqrt_recipm1_alphas_cumprod", torch.sqrt(1.0 / alphas_cumprod - 1))
posterior_variance = betas * (1.0 - alphas_cumprod_prev) / (1.0 - alphas_cumprod)
register("posterior_variance", posterior_variance)
register("posterior_log_variance_clipped", torch.log(posterior_variance.clamp(min=1e-20)))
register(
"posterior_mean_coef1",
betas * torch.sqrt(alphas_cumprod_prev) / (1.0 - alphas_cumprod),
)
register(
"posterior_mean_coef2",
(1.0 - alphas_cumprod_prev) * torch.sqrt(alphas) / (1.0 - alphas_cumprod),
)
register(
"loss_weight",
torch.sqrt(alphas) * torch.sqrt(1.0 - alphas_cumprod) / betas / 100,
)
def predict_noise_from_start(self, x_t, t, x0):
return (extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - x0) / extract(
self.sqrt_recipm1_alphas_cumprod, t, x_t.shape
)
def q_posterior(self, x_start, x_t, t):
posterior_mean = (
extract(self.posterior_mean_coef1, t, x_t.shape) * x_start
+ extract(self.posterior_mean_coef2, t, x_t.shape) * x_t
)
posterior_variance = extract(self.posterior_variance, t, x_t.shape)
posterior_log_variance = extract(self.posterior_log_variance_clipped, t, x_t.shape)
return posterior_mean, posterior_variance, posterior_log_variance
def output(self, x, t):
"""Model forward: predict x_0 as trend + season."""
trend, season = self.model(x, t)
return trend + season
def model_predictions(self, x, t):
"""Predict x_0 (NO clamp -- returns are unbounded) and derive noise."""
x_start = self.output(x, t)
pred_noise = self.predict_noise_from_start(x, t, x_start)
return pred_noise, x_start
def p_mean_variance(self, x, t):
"""Posterior mean and variance -- NO clamp on x_start."""
_, x_start = self.model_predictions(x, t)
model_mean, posterior_variance, posterior_log_variance = self.q_posterior(
x_start=x_start, x_t=x, t=t
)
return model_mean, posterior_variance, posterior_log_variance, x_start
def p_sample(self, x, t: int, cond_fn=None, model_kwargs=None):
"""Single DDPM reverse step with optional classifier guidance."""
batched_times = torch.full((x.shape[0],), t, device=x.device, dtype=torch.long)
model_mean, _, model_log_variance, x_start = self.p_mean_variance(x=x, t=batched_times)
noise = torch.randn_like(x) if t > 0 else 0.0
if cond_fn is not None:
model_mean = self.condition_mean(
cond_fn,
model_mean,
model_log_variance,
x,
t=batched_times,
model_kwargs=model_kwargs,
)
return model_mean + (0.5 * model_log_variance).exp() * noise, x_start
@torch.no_grad()
def sample(self, shape):
"""Full DDPM reverse sampling (all timesteps)."""
img = torch.randn(shape, device=self.betas.device)
for t in reversed(range(self.num_timesteps)):
img, _ = self.p_sample(img, t)
return img
@torch.no_grad()
def fast_sample(self, shape):
"""DDIM-style accelerated sampling."""
batch, total_timesteps = shape[0], self.num_timesteps
sampling_timesteps, eta = self.sampling_timesteps, self.eta
times = torch.linspace(-1, total_timesteps - 1, steps=sampling_timesteps + 1)
times = list(reversed(times.int().tolist()))
time_pairs = list(zip(times[:-1], times[1:], strict=False))
img = torch.randn(shape, device=self.betas.device)
for time, time_next in time_pairs:
time_cond = torch.full((batch,), time, device=self.betas.device, dtype=torch.long)
pred_noise, x_start = self.model_predictions(img, time_cond)
if time_next < 0:
img = x_start
continue
alpha = self.alphas_cumprod[time]
alpha_next = self.alphas_cumprod[time_next]
sigma = eta * ((1 - alpha / alpha_next) * (1 - alpha_next) / (1 - alpha)).sqrt()
c = (1 - alpha_next - sigma**2).sqrt()
noise = torch.randn_like(img)
img = x_start * alpha_next.sqrt() + c * pred_noise + sigma * noise
return img
@torch.no_grad()
def fast_sample_cond(self, shape, cond_fn=None, model_kwargs=None, eta=None):
"""DDIM-style sampling with classifier guidance."""
batch, total_timesteps = shape[0], self.num_timesteps
sampling_timesteps = self.sampling_timesteps
eta = eta if eta is not None else self.eta
times = torch.linspace(-1, total_timesteps - 1, steps=sampling_timesteps + 1)
times = list(reversed(times.int().tolist()))
time_pairs = list(zip(times[:-1], times[1:], strict=False))
img = torch.randn(shape, device=self.betas.device)
for time, time_next in time_pairs:
time_cond = torch.full((batch,), time, device=self.betas.device, dtype=torch.long)
pred_noise, x_start = self.model_predictions(img, time_cond)
if cond_fn is not None:
_, x_start = self.condition_score(
cond_fn, x_start, img, time_cond, model_kwargs=model_kwargs
)
pred_noise = self.predict_noise_from_start(img, time_cond, x_start)
if time_next < 0:
img = x_start
continue
alpha = self.alphas_cumprod[time]
alpha_next = self.alphas_cumprod[time_next]
sigma = eta * ((1 - alpha / alpha_next) * (1 - alpha_next) / (1 - alpha)).sqrt()
c = (1 - alpha_next - sigma**2).sqrt()
noise = torch.randn_like(img)
img = x_start * alpha_next.sqrt() + c * pred_noise + sigma * noise
return img
@torch.no_grad()
def sample_cond(self, shape, cond_fn=None, model_kwargs=None, eta=None):
"""Full DDPM reverse sampling with classifier guidance.
Note: eta parameter is ignored for full DDPM (inherently stochastic).
"""
img = torch.randn(shape, device=self.betas.device)
for t in reversed(range(self.num_timesteps)):
img, _ = self.p_sample(img, t, cond_fn=cond_fn, model_kwargs=model_kwargs)
return img
def generate_mts(self, batch_size=16, model_kwargs=None, cond_fn=None, eta=None):
"""Entry point: generate multivariate time series.
Args:
eta: Stochasticity for DDIM sampling. eta=0 is deterministic, eta=1 is full DDPM.
For conditional sampling, eta>0 adds diversity and prevents mode collapse.
"""
shape = (batch_size, self.seq_length, self.feature_size)
if cond_fn is not None:
sample_fn = self.fast_sample_cond if self.fast_sampling else self.sample_cond
return sample_fn(shape, cond_fn=cond_fn, model_kwargs=model_kwargs, eta=eta)
sample_fn = self.fast_sample if self.fast_sampling else self.sample
return sample_fn(shape)
@property
def loss_fn(self):
if self.loss_type == "l1":
return F.l1_loss
elif self.loss_type == "l2":
return F.mse_loss
raise ValueError(f"invalid loss type {self.loss_type}")
def q_sample(self, x_start, t, noise=None):
"""Forward process: add noise to x_0."""
if noise is None:
noise = torch.randn_like(x_start)
return (
extract(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start
+ extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise
)
def _train_loss(self, x_start, t, target=None, noise=None):
"""Combined time-domain (L1) + frequency-domain (Fourier) loss."""
if noise is None:
noise = torch.randn_like(x_start)
if target is None:
target = x_start
x = self.q_sample(x_start=x_start, t=t, noise=noise)
model_out = self.output(x, t)
train_loss = self.loss_fn(model_out, target, reduction="none")
# Fourier loss: match spectral content
fft1 = torch.fft.fft(model_out.transpose(1, 2), norm="forward")
fft2 = torch.fft.fft(target.transpose(1, 2), norm="forward")
fft1, fft2 = fft1.transpose(1, 2), fft2.transpose(1, 2)
fourier_loss = self.loss_fn(
torch.real(fft1), torch.real(fft2), reduction="none"
) + self.loss_fn(torch.imag(fft1), torch.imag(fft2), reduction="none")
train_loss = train_loss + self.ff_weight * fourier_loss
train_loss = reduce(train_loss, "b ... -> b (...)", "mean")
train_loss = train_loss * extract(self.loss_weight, t, train_loss.shape)
return train_loss.mean()
def forward(self, x, **kwargs):
b, c, n, device = *x.shape, x.device
assert n == self.feature_size, f"expected {self.feature_size} features, got {n}"
t = torch.randint(0, self.num_timesteps, (b,), device=device).long()
return self._train_loss(x_start=x, t=t, **kwargs)
def return_components(self, x, t: int):
"""Return trend, seasonal, residual decomposition for visualization."""
b, c, n, device = *x.shape, x.device
t_tensor = torch.tensor([t]).repeat(b).to(device)
x_noised = self.q_sample(x, t_tensor)
trend, season, residual = self.model(x_noised, t_tensor, return_res=True)
return trend, season, residual, x_noised
def condition_mean(self, cond_fn, mean, log_variance, x, t, model_kwargs=None):
"""Shift mean by σ² · ∇_x log p(y|x) for classifier guidance."""
gradient = cond_fn(x=x, t=t, **(model_kwargs or {}))
return mean.float() + torch.exp(log_variance) * gradient.float()
def condition_score(self, cond_fn, x_start, x, t, model_kwargs=None):
"""Score-based conditioning (Song et al. 2020)."""
alpha_bar = extract(self.alphas_cumprod, t, x.shape)
eps = self.predict_noise_from_start(x, t, x_start)
eps = eps - (1 - alpha_bar).sqrt() * cond_fn(x=x, t=t, **(model_kwargs or {}))
pred_xstart = (
extract(self.sqrt_recip_alphas_cumprod, t, x.shape) * x
- extract(self.sqrt_recipm1_alphas_cumprod, t, x.shape) * eps
)
model_mean, _, _ = self.q_posterior(x_start=pred_xstart, x_t=x, t=t)
return model_mean, pred_xstart
# %% [markdown]
# ## 6. Exponential Moving Average
#
# EMA maintains a shadow copy of model weights that is updated as a running
# average: $\theta_{\text{ema}} \leftarrow \beta \theta_{\text{ema}} + (1-\beta) \theta$.
# Sampling from the EMA model produces smoother, higher-quality outputs.
# %%
class EMA:
"""Simple exponential moving average of model parameters."""
def __init__(self, model, decay=0.995, update_every=10):
self.decay = decay
self.update_every = update_every
self.step = 0
self.ema_model = deepcopy(model)
self.ema_model.eval()
for p in self.ema_model.parameters():
p.requires_grad_(False)
def update(self, model):
self.step += 1
if self.step % self.update_every != 0:
return
with torch.no_grad():
for ema_p, model_p in zip(
self.ema_model.parameters(), model.parameters(), strict=False
):
ema_p.data.mul_(self.decay).add_(model_p.data, alpha=1.0 - self.decay)
def to(self, device):
self.ema_model = self.ema_model.to(device)
return self
# %% [markdown]
# ## 7. Training
#
# Training follows the standard diffusion objective: sample a timestep $t$,
# add noise to create $x_t$, predict $\hat{x}_0$, and minimize the combined
# time-domain + Fourier loss. We use gradient clipping, warmup scheduler,
# gradient accumulation, and EMA for stable convergence.
# %%
# Normalize data before training
scaler = StandardScaler()
scaler.fit(returns) # Fit on raw returns, not sequences
seq_shape = sequences.shape
sequences_flat = sequences.reshape(-1, seq_shape[-1])
sequences_norm = scaler.transform(sequences_flat).reshape(seq_shape).astype(np.float32)
print(f"Normalized: mean={sequences_norm.mean():.4f}, std={sequences_norm.std():.4f}")
# %%
# Initialize model
diffusion_model = DiffusionTS(
seq_length=CONFIG["seq_length"],
feature_size=n_assets,
n_layer_enc=CONFIG["n_layer_enc"],
n_layer_dec=CONFIG["n_layer_dec"],
d_model=CONFIG["d_model"],
timesteps=CONFIG["timesteps"],
sampling_timesteps=CONFIG["sampling_timesteps"],
eta=CONFIG["eta"],
loss_type=CONFIG["loss_type"],
n_heads=CONFIG["n_heads"],
).to(device)
n_params = sum(p.numel() for p in diffusion_model.parameters())
print(f"Model parameters: {n_params:,}")
print(f"DDIM: {CONFIG['timesteps']} training steps → {CONFIG['sampling_timesteps']} sampling steps")
# %% [markdown]
# ### Checkpoint Loading
#
# If a trained model exists and `RETRAIN=False`, we skip training and load
# the saved weights. This allows iterating on evaluation/visualization
# without retraining.
# %%
# Check for existing checkpoint
checkpoint_exists = CHECKPOINT_PATH.exists()
training_losses = []
if checkpoint_exists and not RETRAIN:
print(f"Loading checkpoint from {CHECKPOINT_PATH}")
checkpoint = torch.load(CHECKPOINT_PATH, map_location=device, weights_only=False)
diffusion_model.load_state_dict(checkpoint["model_state"])
scaler = StandardScaler()
scaler.mean_ = np.array(checkpoint["scaler_mean"])
scaler.scale_ = np.array(checkpoint["scaler_scale"])
training_losses = checkpoint.get("training_losses", [])
print(f"Loaded model trained for {len(training_losses)} epochs")
# Create EMA wrapper with loaded weights
ema = EMA(diffusion_model, decay=CONFIG["ema_decay"]).to(device)
ema.ema_model.load_state_dict(checkpoint["ema_state"])
SKIP_TRAINING = True
else:
if checkpoint_exists:
print("RETRAIN=True: Ignoring existing checkpoint")
else:
print(f"No checkpoint found at {CHECKPOINT_PATH}")
SKIP_TRAINING = False
# %%
# Create tensor data and utilities (needed for visualization/classifier even when loading from checkpoint)
tensor_data = torch.FloatTensor(sequences_norm).to(device)
grad_accum = CONFIG["gradient_accumulate_every"]
def cycle(dl):
"""Infinite iterator over a dataloader."""
while True:
yield from dl
# %% [markdown]
# ### Training (or Loading from Checkpoint)
# %%
# Training setup: optimizer, scheduler, dataloader
if not SKIP_TRAINING:
ema = EMA(diffusion_model, decay=CONFIG["ema_decay"]).to(device)
optimizer = torch.optim.Adam(diffusion_model.parameters(), lr=CONFIG["lr"], betas=(0.9, 0.96))
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
optimizer, factor=0.5, patience=SCHEDULER_PATIENCE, min_lr=1e-5, threshold=0.1
)
warmup_steps = CONFIG["warmup_steps"]
warmup_lr = CONFIG["warmup_lr"]
base_lr = CONFIG["lr"]
dataset = TensorDataset(tensor_data)
dataloader = DataLoader(dataset, batch_size=CONFIG["batch_size"], shuffle=True, drop_last=True)
data_iter = cycle(dataloader)
print(f"Training samples: {len(dataset):,}")
print(f"Effective batch size: {CONFIG['batch_size'] * grad_accum}")
# %%
# Training loop
if not SKIP_TRAINING:
print(f"Training Diffusion-TS for {CONFIG['epochs']} epochs...")
training_losses = []
for step in range(CONFIG["epochs"]):
diffusion_model.train()
total_loss = 0.0
# Warmup LR
if step < warmup_steps:
lr = base_lr + (warmup_lr - base_lr) * step / max(warmup_steps, 1)
for pg in optimizer.param_groups:
pg["lr"] = lr
# Gradient accumulation
for _ in range(grad_accum):
(batch,) = next(data_iter)
loss = diffusion_model(batch, target=batch)
loss = loss / grad_accum
loss.backward()
total_loss += loss.item()
torch.nn.utils.clip_grad_norm_(diffusion_model.parameters(), 1.0)
optimizer.step()
if step >= warmup_steps:
scheduler.step(total_loss)
optimizer.zero_grad()
ema.update(diffusion_model)
training_losses.append(total_loss)
if (step + 1) % 100 == 0 or step == 0:
current_lr = optimizer.param_groups[0]["lr"]
print(
f" Epoch {step + 1}/{CONFIG['epochs']}: Loss = {total_loss:.6f}, LR = {current_lr:.2e}",
flush=True,
)
print(f"Training complete. Final loss: {training_losses[-1]:.6f}")
# %%
# Save checkpoint
if not SKIP_TRAINING:
checkpoint = {
"model_state": diffusion_model.state_dict(),
"ema_state": ema.ema_model.state_dict(),
"scaler_mean": scaler.mean_.tolist(),
"scaler_scale": scaler.scale_.tolist(),
"training_losses": training_losses,
"config": CONFIG,
}
torch.save(checkpoint, CHECKPOINT_PATH)
print(f"Saved checkpoint to {CHECKPOINT_PATH}")
# %% [markdown]
# ### Training Progress
# %%
if training_losses:
fig, ax = plt.subplots(figsize=(8, 4), constrained_layout=True)
ax.plot(training_losses, linewidth=1)
ax.set_yscale("log")
ax.set_xlabel("Epoch")
ax.set_ylabel("Loss (L1 + Fourier)")
ax.set_title("Diffusion-TS training loss by epoch")
show_with_alt(
fig,
"Training loss against epoch on a logarithmic vertical axis. The curve falls "
"very steeply over the first few hundred epochs, then flattens into a narrow "
"noisy band that drifts down slightly and stays there to the last epoch.",
)
else:
print("No training losses available (loaded from checkpoint)")
# %% [markdown]
# ## 8. Generate Synthetic Sequences
#
# We sample from the EMA model using DDIM fast sampling (50 steps instead
# of 500). The samples are generated in normalized space and then
# denormalized back to original return scale.
# %%
# N_SYNTHETIC is set in the parameters cell above
print(
f"Generating {N_SYNTHETIC} synthetic sequences via DDIM ({CONFIG['sampling_timesteps']} steps)..."
)
synthetic_norm = ema.ema_model.generate_mts(batch_size=N_SYNTHETIC).detach().cpu().numpy()
# Diagnostic: compare normalized variance
print("\nNormalized space (before scaling):")
print(f" Training std: {sequences_norm.std():.4f}")
print(f" Synthetic std: {synthetic_norm.std():.4f}")
variance_ratio = synthetic_norm.std() / sequences_norm.std()
print(f" Ratio: {variance_ratio:.2%}")
# Variance scaling: Diffusion-TS with trend+seasonal decomposition tends to underestimate
# variance. We scale outputs to match training distribution variance.
# This scale_factor is also applied to regime-conditional samples below.
if variance_ratio < 0.9:
VARIANCE_SCALE_FACTOR = sequences_norm.std() / synthetic_norm.std()
synthetic_norm = synthetic_norm * VARIANCE_SCALE_FACTOR
print(f"\nApplied variance scaling: {VARIANCE_SCALE_FACTOR:.3f}x")
print(f" Scaled std: {synthetic_norm.std():.4f}")
else:
VARIANCE_SCALE_FACTOR = 1.0
# Denormalize
syn_shape = synthetic_norm.shape
synthetic_flat = synthetic_norm.reshape(-1, syn_shape[-1])
synthetic_sequences = scaler.inverse_transform(synthetic_flat).reshape(syn_shape).astype(np.float32)
print("\nDenormalized (return space):")
print(f" Synthetic: mean={synthetic_sequences.mean():.6f}, std={synthetic_sequences.std():.6f}")
print(f" Real: mean={sequences.mean():.6f}, std={sequences.std():.6f}")
# %% [markdown]
# ## 9. Unconditional Evaluation
#
# We evaluate the generated data on three axes:
#
# ### Statistical Tests
#
# - **Kolmogorov-Smirnov (KS) test**: Measures the maximum distance between
# two empirical CDFs. Values range from 0 for identical distributions to 1 for
# completely separated ones, so a small value indicates a close marginal fit.
#
# - **Correlation error**: Mean absolute difference between real and synthetic
# cross-asset correlation matrices. Captures whether the model learned
# dependence structure (e.g., sector correlations).
#
# - **Autocorrelation error**: Compares lag-1 autocorrelation. Returns have
# near-zero AC (weak form efficiency) but volatility clusters (squared returns
# have positive AC). Good generators preserve these stylized facts.
#
# ### Visual Comparison
#
# - **PCA/t-SNE**: Project high-dimensional sequences to 2D. Real and synthetic
# distributions should overlap if the generator captures the data manifold.
#
# ### Utility (TSTR)
#
# - **Train-Synthetic-Test-Real**: Train a classifier on synthetic data, test
# on real data. High accuracy means synthetic data is useful for downstream
# ML tasks -- the ultimate practical validation.
# %%
def evaluate_statistics(real_data: np.ndarray, synthetic_data: np.ndarray) -> dict:
"""Compare distributional properties of real and synthetic data."""
n_assets = real_data.shape[2]
real_flat = real_data.reshape(-1, n_assets)
syn_flat = synthetic_data.reshape(-1, n_assets)
# KS test per asset
ks_stats = [stats.ks_2samp(real_flat[:, i], syn_flat[:, i])[0] for i in range(n_assets)]
# Correlation matrix comparison
real_corr = np.corrcoef(real_flat.T)
syn_corr = np.corrcoef(syn_flat.T)
corr_error = np.mean(np.abs(real_corr - syn_corr))
# Autocorrelation (lag-1) per asset
def autocorr(x, lag=1):
return np.corrcoef(x[:-lag], x[lag:])[0, 1]
real_ac = [autocorr(real_flat[:, i]) for i in range(n_assets)]
syn_ac = [autocorr(syn_flat[:, i]) for i in range(n_assets)]
ac_error = np.mean(np.abs(np.array(real_ac) - np.array(syn_ac)))
# Every figure below averages over assets. Carry the spread as well, so a mean
# that hides one badly-fitted asset is visible as one (standard C18).
return {
"n_assets": n_assets,
"mean_ks_statistic": np.mean(ks_stats),
"worst_ks_statistic": np.max(ks_stats),
"best_ks_statistic": np.min(ks_stats),
"mean_error": np.mean(np.abs(real_flat.mean(0) - syn_flat.mean(0))),
"std_error": np.mean(np.abs(real_flat.std(0) - syn_flat.std(0))),
"correlation_error": corr_error,
"autocorrelation_error": ac_error,
"worst_autocorrelation_error": np.max(np.abs(np.array(real_ac) - np.array(syn_ac))),
}
stats_results = evaluate_statistics(sequences, synthetic_sequences)
print("\n=== Statistical Evaluation ===")
for key, value in stats_results.items():
print(f" {key}: {value:.4f}" if isinstance(value, float) else f" {key}: {value}")
# %% [markdown]
# **Interpretation**: a low KS statistic means the marginal distribution for an asset
# is well matched. Correlation error tests whether cross-asset dependence survived, and
# autocorrelation error whether the weak serial dependence of daily returns did.
#
# Read each mean against the range printed beside it, not on its own. Five of the nine
# printed numbers are means. Four of those - the KS statistic, the mean error, the
# standard-deviation error and the autocorrelation error - average over assets, while
# the correlation error averages the absolute difference between corresponding entries
# of the real and synthetic correlation matrices, so it averages over asset pairs. The
# other four are not means: `n_assets` counts the assets, and the two extreme KS values
# and the largest autocorrelation error are each one asset's.
#
# All of these errors are absolute, so a small mean cannot come from large errors
# cancelling; it comes from many small ones diluting a few large ones. The extreme
# per-asset KS values say whether that happened: a maximum near the mean means the fit
# is even across assets, and one far above it means the mean describes the assets the
# model handles and conceals the one it does not. The maximum per-asset autocorrelation
# error reads the same way.
# %%
fig = plot_fidelity_comparison(
sequences,
synthetic_sequences,
title="Diffusion-TS: Real vs Synthetic Distribution",
n_samples=500,
flatten_method="flatten", # Flatten all timesteps for full sequence comparison
)
show_with_alt(
fig,
"Two scatter panels comparing real and synthetic sequences. In the PCA projection "
"both sets sit in one dense cluster at the origin, with scattered real outliers far "
"from it in several directions and a few synthetic ones closer in. In the t-SNE "
"projection the two sets are interleaved across the whole area with no region "
"belonging to only one of them.",
)
# %% [markdown]
# **Interpretation**: Overlapping PCA/t-SNE point clouds confirm that synthetic
# sequences occupy the same region of feature space as real data. Gaps or
# isolated clusters would indicate missing regimes.
# %% [markdown]
# ### TSTR Evaluation
#
# We evaluate downstream utility via **extreme-move classification**: predict
# whether the next-day absolute return exceeds the 90th percentile.
#
# **Important context**: this is a heavily imbalanced task. The threshold is a high
# percentile of absolute *training* returns, so only that tail fraction of the training
# window is positive by construction. Applied out-of-sample to a calmer holdout period,
# fewer test days still clear the bar, which pushes the accuracy of a classifier that
# never predicts the positive class very high. Both rates are printed below. Read the
# **TSTR ratio** (synthetic over real), not raw accuracy.
# %%
def tstr_evaluation(train_data, holdout_data, synthetic_data):
"""TSTR on extreme-move classification (90th percentile threshold)."""
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import precision_recall_fscore_support
train_returns = train_data[:, -1, 0]
threshold = np.percentile(np.abs(train_returns), 90)
X_train_real = train_data[:, :-1, :].reshape(len(train_data), -1)
y_train_real = (np.abs(train_returns) > threshold).astype(int)
X_train_syn = synthetic_data[:, :-1, :].reshape(len(synthetic_data), -1)
y_train_syn = (np.abs(synthetic_data[:, -1, 0]) > threshold).astype(int)
X_test = holdout_data[:, :-1, :].reshape(len(holdout_data), -1)
y_test = (np.abs(holdout_data[:, -1, 0]) > threshold).astype(int)
scaler_r = StandardScaler()
X_tr_s = scaler_r.fit_transform(X_train_real)
X_te_r = scaler_r.transform(X_test)
scaler_s = StandardScaler()
X_ts_s = scaler_s.fit_transform(X_train_syn)
X_te_s = scaler_s.transform(X_test)
model_r = LogisticRegression(max_iter=1000)
model_r.fit(X_tr_s, y_train_real)
acc_real = model_r.score(X_te_r, y_test)
y_pred_real = model_r.predict(X_te_r)
prec_r, rec_r, f1_r, _ = precision_recall_fscore_support(
y_test, y_pred_real, average="binary", zero_division=0
)
if len(np.unique(y_train_syn)) < 2:
acc_syn = (y_test == int(y_train_syn.mean() > 0.5)).mean()
prec_s, rec_s, f1_s = 0, 0, 0
else:
model_s = LogisticRegression(max_iter=1000)
model_s.fit(X_ts_s, y_train_syn)
acc_syn = model_s.score(X_te_s, y_test)
y_pred_syn = model_s.predict(X_te_s)
prec_s, rec_s, f1_s, _ = precision_recall_fscore_support(
y_test, y_pred_syn, average="binary", zero_division=0
)
return {
"accuracy_real": acc_real,
"accuracy_synthetic": acc_syn,
"precision_real": prec_r,
"precision_synthetic": prec_s,
Exibido na íntegra, com atribuição conforme a licença da fonte. Licença: MIT
Este resumo foi escrito pelo agente de pesquisa da Stratmill com base no original; não é uma cópia da fonte.