Skip to content
All library documents

Interpretable Diffusion Models for Synthetic Financial Returns

Code Machine Learning for Trading

Summary

This notebook describes adapting Diffusion-TS to generate synthetic ETF return sequences. Its denoiser predicts the clean sequence as a sum of polynomial trend, Fourier seasonal, and residual components, while a combined time-domain and Fourier-domain loss encourages both pointwise and spectral fidelity. A DDIM sampler reduces the number of reverse steps used for generation. A separate classifier trained on noised sequences provides guidance for generating samples labeled as low or high volatility regimes.

The notebook evaluates unconditional and regime-conditional samples against historical data using distributional, correlation, autocorrelation, and downstream prediction measures. It reports that lower guidance strength and higher sampling temperature are used for the minority high-volatility regime to avoid collapsing its diversity. The model is adapted for unbounded standardized returns without image-style clipping. Results depend on the simple HMM regime labels and tuned guidance settings; the notebook also notes the computational burden of a multi-asset model and does not establish that synthetic sequences capture every market property needed for trading.

Key ideas

  • Diffusion-TS predicts the clean return sequence and decomposes it into trend, seasonal, and residual components.
  • A Fourier-domain loss is used to encourage preservation of spectral and autocorrelation structure.
  • DDIM sampling reduces the reverse-process steps used to generate sequences.
  • A classifier trained on noised data guides generation toward volatility regimes.
  • Regime guidance settings affect both target fidelity and sample diversity, especially for a minority regime.

Tags

Full text
# 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,
     

Shown in full with attribution under the source's licence. Licence: MIT

This summary was written by Stratmill's research agent from the original; it is not a copy of the source.