"""
时序 Transformer + 注意力权重可视化
------------------------------------------------
演示「注意力 = 加权平均」思想在时间序列预测中的应用：
  1. 用带双周期的正弦序列模拟真实时序
  2. 训练一个单层 Transformer 预测下一个点
  3. 提取 self-attention 权重，画出热力图 + 预测时的权重分配柱状图
     —— 让你亲眼看到模型预测时"给了哪些历史时刻多少权重"

by 小小叶 🍃 · OpenClaw
"""
import torch
import torch.nn as nn
import numpy as np
import math
import matplotlib.pyplot as plt

# ============ 1. 造数据：带周期的时间序列 ============
def make_data(seq_len=48, n=3000):
    """seq_len=48 个历史点 → 预测下 1 个点。
       双周期正弦，模拟'短周期+长周期'的真实时序"""
    t = np.linspace(0, 200, n)
    series = np.sin(t) + 0.5 * np.sin(0.25 * t) + 0.1 * np.random.randn(n)
    X, Y = [], []
    for i in range(len(series) - seq_len - 1):
        X.append(series[i:i + seq_len])
        Y.append(series[i + seq_len])
    X = torch.tensor(np.array(X), dtype=torch.float32).unsqueeze(-1)  # [N,seq,1]
    Y = torch.tensor(np.array(Y), dtype=torch.float32).unsqueeze(-1)  # [N,1]
    return X, Y, series

# ============ 2. 位置编码 ============
class PositionalEncoding(nn.Module):
    def __init__(self, d_model, max_len=500):
        super().__init__()
        pe = torch.zeros(max_len, d_model)
        pos = torch.arange(max_len).unsqueeze(1).float()
        div = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
        pe[:, 0::2] = torch.sin(pos * div)
        pe[:, 1::2] = torch.cos(pos * div)
        self.register_buffer('pe', pe.unsqueeze(0))

    def forward(self, x):
        return x + self.pe[:, :x.size(1)]

# ============ 3. 自定义 Attention Block（便于提取注意力权重）============
class AttnBlock(nn.Module):
    def __init__(self, d_model, nhead):
        super().__init__()
        self.attn = nn.MultiheadAttention(d_model, nhead, batch_first=True)
        self.norm1 = nn.LayerNorm(d_model)
        self.ff = nn.Sequential(nn.Linear(d_model, 128), nn.ReLU(), nn.Linear(128, d_model))
        self.norm2 = nn.LayerNorm(d_model)
        self.attn_weights = None

    def forward(self, x):
        out, w = self.attn(x, x, x, need_weights=True, average_attn_weights=True)
        self.attn_weights = w.detach()  # [B, seq, seq]
        x = self.norm1(x + out)
        x = self.norm2(x + self.ff(x))
        return x

class TSTransformer(nn.Module):
    def __init__(self, d_model=64, nhead=4):
        super().__init__()
        self.proj = nn.Linear(1, d_model)
        self.pos = PositionalEncoding(d_model)
        self.block = AttnBlock(d_model, nhead)
        self.head = nn.Linear(d_model, 1)

    def forward(self, x):
        x = self.pos(self.proj(x))
        x = self.block(x)
        return self.head(x[:, -1, :])  # 用最后时刻预测

# ============ 4. 训练 ============
def main():
    X, Y, series = make_data()
    ntr = int(len(X) * 0.8)
    Xtr, Ytr, Xte, Yte = X[:ntr], Y[:ntr], X[ntr:], Y[ntr:]

    model = TSTransformer()
    opt = torch.optim.Adam(model.parameters(), lr=1e-3)
    loss_fn = nn.MSELoss()

    for ep in range(40):
        model.train()
        perm = torch.randperm(len(Xtr))
        for i in range(0, len(Xtr), 64):
            idx = perm[i:i + 64]
            opt.zero_grad()
            loss = loss_fn(model(Xtr[idx]), Ytr[idx])
            loss.backward()
            opt.step()
        if (ep + 1) % 10 == 0:
            model.eval()
            with torch.no_grad():
                te = loss_fn(model(Xte), Yte).item()
            print(f"Epoch {ep+1} | test MSE {te:.4f}")

    # ============ 5. 提取并可视化注意力权重 ============
    model.eval()
    sample = Xte[0:1]
    with torch.no_grad():
        pred = model(sample)
    attn = model.block.attn_weights[0].numpy()  # [seq, seq]

    # 图1：完整注意力热力图
    plt.figure(figsize=(8, 6))
    plt.imshow(attn, cmap='viridis', aspect='auto')
    plt.colorbar(label='attention weight')
    plt.xlabel('Key position (history step being attended)')
    plt.ylabel('Query position (step issuing attention)')
    plt.title('Self-Attention weight heatmap')
    plt.tight_layout()
    plt.savefig('attn_heatmap.png', dpi=120)

    # 图2：预测时（最后时刻）的加权平均分配
    last_q_weights = attn[-1]
    plt.figure(figsize=(10, 4))
    plt.bar(range(len(last_q_weights)), last_q_weights, color='teal')
    plt.xlabel('history step (0=oldest, 47=latest)')
    plt.ylabel('weight')
    plt.title('Weighted-average allocation over history when predicting next point')
    plt.tight_layout()
    plt.savefig('attn_laststep.png', dpi=120)

    print("\nTop5 history steps attended at last position:")
    top5 = last_q_weights.argsort()[::-1][:5]
    for idx in top5:
        print(f"  step {idx:2d} ({47-idx} steps ago): weight {last_q_weights[idx]:.3f}")
    print("\nSaved: attn_heatmap.png & attn_laststep.png")

if __name__ == '__main__':
    main()
