해석 가능한 금융 수익률 확산 모델
코드 Machine Learning for Trading
요약
이 노트북은 합성 ETF 수익률 시퀀스를 생성하도록 Diffusion-TS을 적용하는 방법을 설명합니다. 노이즈 제거기는 원래의 깨끗한 시퀀스를 다항식 추세, 푸리에 계절성, 잔차 성분의 합으로 예측하며, 시간 영역과 푸리에 영역의 손실을 결합해 각 지점과 주파수 구조의 충실도를 모두 높입니다. DDIM 샘플러는 생성에 쓰이는 역방향 단계를 줄입니다. 노이즈가 있는 시퀀스로 학습된 별도 분류기는 저변동성 또는 고변동성 국면 레이블이 붙은 표본을 만들도록 유도합니다.
노트북은 분포, 상관관계, 자기상관, 후속 예측 지표를 사용해 무조건부 및 국면 조건부 표본을 과거 데이터와 비교합니다. 다양성 붕괴를 막기 위해 소수인 고변동성 국면에는 더 낮은 유도 강도와 더 높은 샘플링 온도를 사용한다고 보고합니다. 모델은 이미지 생성에서 사용하는 값 클리핑을 적용하지 않고 범위 제한이 없는 표준화 수익률에 맞게 조정됩니다. 결과는 단순한 HMM 국면 레이블과 조정된 유도 설정에 달려 있습니다. 노트북은 다중 자산 모델의 계산 부담도 언급하며 합성 시퀀스가 트레이딩에 필요한 모든 시장 특성을 포착한다고 입증하지 않습니다.
핵심 아이디어
- Diffusion-TS은 노이즈가 제거된 수익률 시퀀스를 예측하고 이를 추세, 계절성, 잔차로 분해합니다.
- 푸리에 영역 손실로 주파수와 자기상관 구조를 보존하도록 유도합니다.
- DDIM 샘플링은 시퀀스 생성에 필요한 역방향 과정의 단계를 줄입니다.
- 노이즈가 있는 데이터로 학습한 분류기가 변동성 국면에 맞는 생성을 유도합니다.
- 국면 유도 설정은 목표 충실도와 표본 다양성 모두에 영향을 주며, 특히 소수 국면에서 그렇습니다.
태그
전문
# 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,
출처의 라이선스에 따라 출처를 표시하고 전문을 공개합니다. 라이선스: MIT
이 요약은 원문을 바탕으로 Stratmill의 리서치 에이전트가 작성했으며, 원문을 복사한 것이 아닙니다.