Pular para o conteúdo
Todos os documentos da biblioteca

Modelos de difusão interpretáveis para dados financeiros condicionados a regimes

Notebook Machine Learning for Trading

Resumo

Este notebook adapta o Diffusion-TS para gerar sequências sintéticas de retornos diários de ETF. O denoiser prevê a série original e decompõe essa previsão em uma tendência polinomial, componentes de Fourier selecionados e um resíduo. Um objetivo combinado nos domínios do tempo e da frequência busca preservar tanto os valores dos retornos quanto a estrutura espectral. A configuração usa sequências sobrepostas e um holdout temporal; em seguida, adiciona um classificador treinado em sequências com ruído para orientar a amostragem a regimes de volatilidade baixa ou alta identificados por um modelo oculto de Markov gaussiano.

O notebook descreve a amostragem DDIM como uma forma de reduzir as etapas reversas mantendo uma qualidade útil e avalia amostras incondicionais e condicionadas a regimes com verificações distributivas e preditivas. Também apresenta escolhas práticas de adaptação: sem limitar a saída como em imagens para retornos padronizados e com configurações de orientação diferentes para os dois regimes. São escolhas de modelagem, não evidências de superioridade geral. Os resultados dependem dos rótulos dos regimes e do ajuste da orientação; o HMM é uma fonte simples de rótulos, e ampliar para mais ativos aumenta o custo computacional. A fidelidade sintética não estabelece que uma estratégia de trading terá bom desempenho.

Ideias principais

  • O denoiser prevê a sequência original de retornos como componentes interpretáveis de tendência, sazonalidade de Fourier e resíduo.
  • Uma função de perda no domínio da frequência incentiva as sequências geradas a preservar características espectrais junto com o ajuste no domínio do tempo.
  • Um classificador treinado em sequências com ruído orienta a geração para regimes de volatilidade especificados.
  • A amostragem DDIM reduz o número de etapas reversas usadas na geração.
  • A qualidade das amostras condicionais depende da precisão dos rótulos de regime e das configurações de orientação ajustadas para cada regime.

Tags

Texto completo
# 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_sy

Exibido na íntegra, com atribuição conforme a licença da fonte. Licença: MIT

Este resumo foi escrito pelo agente de pesquisa da Stratmill com base no original; não é uma cópia da fonte.