Modèles de diffusion interprétables pour les données financières conditionnées par régime
Résumé
Ce notebook adapte Diffusion-TS pour générer des séquences synthétiques de rendements quotidiens de ETF. Son débruiteur prédit la série d’origine et décompose cette prédiction en une tendance polynomiale, des composantes de Fourier sélectionnées et un résidu. Une fonction objectif combinant les domaines temporel et fréquentiel vise à préserver à la fois les valeurs des rendements et la structure spectrale. La configuration utilise des séquences chevauchantes et un échantillon de réserve temporel, puis ajoute un classifieur entraîné sur des séquences bruitées pour guider l’échantillonnage vers des régimes de volatilité faible ou élevée identifiés par un modèle de Markov caché gaussien.
Le notebook présente l’échantillonnage DDIM comme un moyen de réduire le nombre d’étapes inverses tout en conservant une qualité utile, et évalue les échantillons inconditionnels et conditionnés par régime à l’aide de contrôles distributionnels et prédictifs. Il décrit aussi des choix pratiques d’adaptation : aucune limitation de la sortie de type image pour les rendements standardisés, et des réglages de guidage différents pour les deux régimes. Ce sont des choix de modélisation, et non des preuves de supériorité générale. Les résultats dépendent des étiquettes de régime et du réglage du guidage ; la source d’étiquettes HMM est simple, et le passage à un plus grand nombre d’actifs augmente le coût de calcul. La fidélité des données synthétiques ne démontre pas qu’une stratégie de trading sera performante.
Idées clés
- Le débruiteur prédit la séquence de rendements d’origine sous forme de composantes interprétables : tendance, saisonnalité de Fourier et résidu.
- Une fonction de perte dans le domaine fréquentiel incite les séquences générées à conserver leurs caractéristiques spectrales en plus de leur ajustement temporel.
- Un classifieur entraîné sur des séquences bruitées oriente la génération vers des régimes de volatilité spécifiés.
- L’échantillonnage DDIM réduit le nombre d’étapes inverses utilisées pour la génération.
- La qualité des échantillons conditionnels dépend de la précision des étiquettes de régime et des réglages de guidage ajustés à chaque régime.
Étiquettes
Texte intégral
# Diffusion-TS: Interpretable Diffusion with Conditional Generation
# 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.
```python
"""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
```
## 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.
```python
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))
```
## 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.
```python
# 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
```
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.
```python
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)
}
```
```python
# 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}")
```
## 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.
```python
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
```
### Create Overlapping Sequences
Sliding windows maximize sample count from limited financial data.
```python
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
```
```python
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}")
```
## 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.
```python
# 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)))
```
```python
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)
```
### 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.
```python
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
```
### 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.
```python
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)])
```
```python
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)
```
```python
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)
```
## 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.
```python
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)
```
### 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.
```python
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
```
## 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.
```python
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))
```
### 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.
```python
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))
```
### 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).
```python
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
```
```python
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
```
```python
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
```
```python
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
```
### 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.
```python
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
```
## 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.
```python
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)
```
```python
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
```
## 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.
```python
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
```
## 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.
```python
# 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}")
```
```python
# 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")
```
### 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.
```python
# 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
```
```python
# 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
```
### Training (or Loading from Checkpoint)
```python
# 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}")
```
```python
# 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}")
```
```python
# 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}")
```
### Training Progress
```python
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)")
```
## 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.
```python
# 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}")
```
## 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.
```python
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}")
```
**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.
```python
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.",
)
```
**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.
### 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.
```python
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,
"recall_real": rec_r,
"recall_synthetic": rec_s,
"f1_real": f1_r,
"f1_synthetic": f1_s,
"tstr_ratio": acc_syn / acc_real if acc_real > 0 else 0,
"positive_rate": y_test.mean(),
"baseline_accuracy": 1 - y_test.mean(),
"n_test_samples": len(y_test),
}
tstr_results = tstr_evaluation(sequences, holdout_sequences, synthetic_sequences)
print("\n=== TSTR Evaluation: Extreme-Move Classification ===")
print(" Task: Predict |return| > 90th percentile")
print(
f" Test samples: {tstr_results['n_test_samples']:,} ({tstr_results['positive_rate']:.1%} positive)"
)
print(f" Naive baseline: {tstr_results['baseline_accuracy']:.1%} (always predict 'normal')")
print()
print(f" {'Metric':<12} {'Real-Trained':>14} {'Synth-Trained':>14}")
print(f" {'-' * 42}")
print(
f" {'Accuracy':<12} {tstr_results['accuracy_real']:>14.1%} {tstr_results['accuracy_syReproduit dans son intégralité avec attribution, conformément à la licence de la source. Licence: MIT
Ce résumé a été rédigé par l’agent de recherche de Stratmill à partir de la source originale ; il n’en est pas une copie.