-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmodel.py
More file actions
713 lines (627 loc) · 32.5 KB
/
Copy pathmodel.py
File metadata and controls
713 lines (627 loc) · 32.5 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
from typing import Optional
# 移除加性位置编码,改用 RoPE 在时间注意力中对 Q/K 做相对位置旋转
def _build_rope_cache(seq_len: int, head_dim: int, device: torch.device, base: float = 10000.0) -> tuple[torch.Tensor, torch.Tensor]:
"""生成 RoPE 所需的 cos/sin 表
Returns:
cos: [seq_len, head_dim//2]
sin: [seq_len, head_dim//2]
"""
half_dim = head_dim // 2
inv_freq = 1.0 / (base ** (torch.arange(0, half_dim, device=device).float() / half_dim)) # [half_dim]
t = torch.arange(seq_len, device=device).float() # [seq_len]
angles = torch.ger(t, inv_freq) # [seq_len, half_dim]
return torch.cos(angles), torch.sin(angles)
def _apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
"""对最后一维 head_dim 应用 RoPE 旋转。
x: [B, H, T, D]
cos/sin: [T, D//2]
"""
B, H, T, D = x.shape
half = D // 2
x_pair = x.view(B, H, T, half, 2)
x1 = x_pair[..., 0]
x2 = x_pair[..., 1]
cos = cos.unsqueeze(0).unsqueeze(0) # [1,1,T,half]
sin = sin.unsqueeze(0).unsqueeze(0) # [1,1,T,half]
y1 = x1 * cos - x2 * sin
y2 = x2 * cos + x1 * sin
y = torch.stack([y1, y2], dim=-1).reshape(B, H, T, D)
return y
class SpacetimeMultiHeadAttention(nn.Module):
"""时空联合多头注意力机制(时间注意力使用 RoPE)"""
def __init__(self, d_model: int, num_heads: int, seq_len: int, dropout: float = 0.1):
super().__init__()
assert d_model % num_heads == 0
self.d_model = d_model
self.num_heads = num_heads
self.d_k = d_model // num_heads
assert self.d_k % 2 == 0, "RoPE 需要每头维度 d_k 为偶数"
self.seq_len = int(seq_len)
# 调试/可视化:是否保存注意力(默认关闭,避免额外内存开销)
self.save_attention = False
self.last_temporal_cls_attn = None # [B, N, H, T]
self.last_spatial_cls_attn = None # [B, H, N, N]
# 可视化所用的CLS数量(由上层注入)
self.k_vis = 6
# 时间注意力
self.temporal_q = nn.Linear(d_model, d_model, bias=False)
self.temporal_k = nn.Linear(d_model, d_model, bias=False)
self.temporal_v = nn.Linear(d_model, d_model, bias=False)
self.temporal_dropout = nn.Dropout(0.3)
self.spatial_proj = nn.Linear(d_model, 2 * d_model)
self.spatial_activation = nn.SiLU()
self.spatial_proj_back = nn.Linear(2 * d_model, d_model)
self.output_proj = nn.Linear(d_model, d_model)
self.dropout = nn.Dropout(dropout)
self.layer_norm = nn.LayerNorm(d_model)
# 预构建并缓存 RoPE cos/sin 表(注册为 buffer,随设备/精度自动迁移)
half_dim = self.d_k // 2
inv_freq = 1.0 / (10000.0 ** (torch.arange(0, half_dim).float() / half_dim)) # [half_dim]
t = torch.arange(self.seq_len).float() # [T]
angles = torch.outer(t, inv_freq) # [T, half_dim]
rope_cos = torch.cos(angles)
rope_sin = torch.sin(angles)
self.register_buffer("rope_cos", rope_cos, persistent=False)
self.register_buffer("rope_sin", rope_sin, persistent=False)
def temporal_attention(self, x: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor:
"""
时间注意力:处理每只股票内部的时序依赖(使用 RoPE 相对位置)
Args:
x: [batch_size, num_stocks, seq_len, d_model]
"""
batch_size, num_stocks, seq_len, d_model = x.shape
# 重塑为 [batch_size * num_stocks, seq_len, d_model]
x_reshaped = x.view(-1, seq_len, d_model)
Q = self.temporal_q(x_reshaped).view(-1, seq_len, self.num_heads, self.d_k).transpose(1, 2) # [B*N, H, T, D]
K = self.temporal_k(x_reshaped).view(-1, seq_len, self.num_heads, self.d_k).transpose(1, 2)
V = self.temporal_v(x_reshaped).view(-1, seq_len, self.num_heads, self.d_k).transpose(1, 2)
# RoPE 应用到 Q、K 上(使用缓存的 cos/sin)
# 若当前 seq_len 变化(极少见),做一次性重建
if self.rope_cos.shape[0] != seq_len:
half_dim = self.d_k // 2
inv_freq = 1.0 / (10000.0 ** (torch.arange(0, half_dim, device=Q.device).float() / half_dim))
t = torch.arange(seq_len, device=Q.device).float()
angles = torch.outer(t, inv_freq)
self.rope_cos = torch.cos(angles)
self.rope_sin = torch.sin(angles)
Q = _apply_rope(Q, self.rope_cos, self.rope_sin)
K = _apply_rope(K, self.rope_cos, self.rope_sin)
# Attention scores: [B, N, H, T+K, T+K]
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
# [mask机制] 应用 mask
if mask is not None:
# mask: [B, N, T+K] -> [B*N, 1, 1, T+K] for broadcasting
attn_mask = mask.view(batch_size * num_stocks, seq_len).unsqueeze(1).unsqueeze(2)
scores = scores.masked_fill(attn_mask == 0, float('-inf'))
# 可视化:保存CLS token对所有时间步的注意力
if self.save_attention:
try:
# attention_weights: [B*N, H, T, T],此处 T=K+T_raw
k_vis = int(getattr(self, 'k_vis', 6))
k_vis = max(1, min(k_vis, seq_len))
cls_rows = scores[:, :, :k_vis, :].mean(dim=2) # [B*N, H, T]
cls_row = cls_rows.view(batch_size, num_stocks, self.num_heads, seq_len)
# 移到CPU以节省显存,并断开梯度
self.last_temporal_cls_attn = cls_row.detach().float().cpu()
except Exception:
self.last_temporal_cls_attn = None
attention_weights = F.softmax(scores, dim=-1)
attention_weights = self.dropout(attention_weights)
context = torch.matmul(attention_weights, V)
context = context.transpose(1, 2).contiguous().view(-1, seq_len, d_model)
# 恢复形状
temporal_output = context.view(batch_size, num_stocks, seq_len, d_model)
# 可视化缓存:对前K个CLS行做平均,得到“CLS聚合”对应的时间注意力(各头对所有时间片)
if self.save_attention:
try:
# attention_weights: [B*N, H, T, T],此处 T=K+T_raw
k_vis = int(getattr(self, 'k_vis', 6))
k_vis = max(1, min(k_vis, seq_len))
cls_rows = attention_weights[:, :, :k_vis, :].mean(dim=2) # [B*N, H, T]
cls_row = cls_rows.view(batch_size, num_stocks, self.num_heads, seq_len)
# 移到CPU以节省显存,并断开梯度
self.last_temporal_cls_attn = cls_row.detach().float().cpu()
except Exception:
self.last_temporal_cls_attn = None
return temporal_output
def spatial_attention(self, x: torch.Tensor, adj_matrix: torch.Tensor) -> torch.Tensor:
"""
空间注意力:处理股票间的关系依赖(向量化所有时间步)
Args:
x: [batch_size, num_stocks, seq_len, d_model]
adj_matrix: [num_stocks, num_stocks] 或 [batch_size, num_stocks, num_stocks]
"""
batch_size, num_stocks, seq_len, d_model = x.shape
if adj_matrix.dim() == 2:
adj = adj_matrix.unsqueeze(0).expand(batch_size, -1, -1) # [B, N, N]
elif adj_matrix.dim() == 3:
adj = adj_matrix
else:
raise ValueError("adj_matrix 维度必须为 2 或 3")
# 展平成 [B*T, N, D],基于残差邻接做一次消息传递
x_btnd = x.permute(0, 2, 1, 3).contiguous().view(batch_size * seq_len, num_stocks, d_model)
adj_bt = adj.unsqueeze(1).repeat(1, seq_len, 1, 1).view(batch_size * seq_len, num_stocks, num_stocks)
context = torch.bmm(adj_bt, x_btnd) # [B*T, N, D]
spatial_output = context.view(batch_size, seq_len, num_stocks, d_model).permute(0, 2, 1, 3).contiguous()
# 禁用空间注意力正则项
self.last_entropy = None
self.last_head_div = None
if self.save_attention:
self.last_spatial_cls_attn = None
return spatial_output
def forward(self, x: torch.Tensor, adj_matrix: torch.Tensor,
mask: Optional[torch.Tensor] = None) -> torch.Tensor:
"""
时空联合注意力
Args:
x: [batch_size, num_stocks, seq_len, d_model]
adj_matrix: [num_stocks, num_stocks]
"""
residual = x
# 分别计算时间和空间注意力
# Pre-LN: 先对输入做 Norm,再进 Attention,结果加回残差(x)
residual_norm = self.layer_norm(x)
temporal_out = self.temporal_attention(residual_norm, mask) # [batch_size, num_stocks, seq_len, d_model]
temporal_out = self.temporal_dropout(temporal_out)
spatial_base = self.spatial_attention(residual_norm, adj_matrix) # [batch_size, num_stocks, seq_len, d_model]
spatial_double = self.spatial_proj(spatial_base)
spatial_double = self.spatial_activation(spatial_double)
spatial_out = self.spatial_proj_back(spatial_double)
fused_output = temporal_out + spatial_out
# 输出投影
output = self.output_proj(fused_output)
# Pre-LN 结构下,输出直接为 residual + output
return x + output
class SpacetimeTransformerBlock(nn.Module):
"""时空联合Transformer块"""
def __init__(self, d_model: int, num_heads: int, d_ff: int, seq_len: int, dropout: float = 0.1,
use_adaln_cond: bool = True):
super().__init__()
self.spacetime_attention = SpacetimeMultiHeadAttention(d_model, num_heads, seq_len, dropout)
self.feed_forward = nn.Sequential(
nn.Linear(d_model, d_ff),
nn.SiLU(),
nn.Dropout(dropout),
nn.Linear(d_ff, d_model)
)
self.layer_norm = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
# 条件调制(AdaLN/FiLM 风格):对归一化后的输出做仿射变换
self.use_adaln_cond = bool(use_adaln_cond)
if self.use_adaln_cond:
self.cond_gamma = nn.Linear(d_model, d_model)
self.cond_beta = nn.Linear(d_model, d_model)
def forward(self, x: torch.Tensor, adj_matrix: torch.Tensor,
mask: Optional[torch.Tensor] = None,
cond: Optional[torch.Tensor] = None) -> torch.Tensor:
# Pre-LN 结构
# 1. 时空注意力 Block (Pre-LN已在Attention模块内部分处理,但通常 Block 级也要统一)
# 由于 SpacetimeMultiHeadAttention 内部已经实现了 Pre-LN (对输入Norm后计算Attn,返回 x + Attn(Norm(x)))
# 所以这里直接调用即可
attn_output = self.spacetime_attention(x, adj_matrix, mask)
# 记录正则项
self.last_entropy = getattr(self.spacetime_attention, 'last_entropy', None)
self.last_head_div = getattr(self.spacetime_attention, 'last_head_div', None)
# 透传可视化缓存(CLS注意力)
self.last_temporal_cls_attn = getattr(self.spacetime_attention, 'last_temporal_cls_attn', None)
self.last_spatial_cls_attn = getattr(self.spacetime_attention, 'last_spatial_cls_attn', None)
# 2. 前馈网络 Block (Pre-LN)
# 输入是 attn_output (即 x + Attn),对其做 Norm 后进 FFN,再加回残差
ff_input_norm = self.layer_norm(attn_output)
ff_output = self.feed_forward(ff_input_norm)
# 残差连接
output = attn_output + ff_output
# 条件仿射调制:output = (ff + res) * (1 + tanh(gamma(cond))) + beta(cond)
# 注意:在Pre-LN中,最后的输出通常不再做Norm,或者做一次Final Norm。
# 这里 AdaLN 作用在 Block 输出上
if self.use_adaln_cond and (cond is not None):
# 不使用try-except,维度错误时立即报错
gamma = torch.tanh(self.cond_gamma(cond)) # [B, N, D]
beta = self.cond_beta(cond) # [B, N, D]
# output是[B, N, T, D],需要unsqueeze T维度进行广播
gamma = gamma.unsqueeze(2) # [B, N, 1, D]
beta = beta.unsqueeze(2) # [B, N, 1, D]
output = output * (1.0 + gamma) + beta
return output
class DynamicGraphLearner(nn.Module):
"""动态图结构学习器(去除固定 per-stock 嵌入,支持可变 num_stocks)"""
def __init__(self, d_model: int, temp: float = 1.0,
adj_activation: str = "softmax",
entmax_alpha: float = 1.5):
super().__init__()
self.d_model = d_model
self.temp = temp
# 稀疏邻接激活
self.adj_activation = str(adj_activation).lower().strip()
self.entmax_alpha = float(entmax_alpha)
self.last_learned_adj: Optional[torch.Tensor] = None
self.last_frobenius: Optional[torch.Tensor] = None
self.last_adj_raw_abs_mean: Optional[torch.Tensor] = None
# Attention Pooling 权重
self.time_pooling_weight = nn.Linear(d_model, 1)
# 关系学习网络(使用 Swish/SiLU)
self.relation_mlp = nn.Sequential(
nn.Linear(d_model * 2, d_model),
nn.SiLU(),
nn.Linear(d_model, 1)
)
# 已移除 top-k/退火机制
def _entmax_bisect(self, z: torch.Tensor, alpha: float = 1.5, dim: int = -1, n_iter: int = 50) -> torch.Tensor:
"""近似 Entmax(alpha)(二分法),返回非负且行和为1的稀疏概率。
参考文献: Peters et al., "Sparse Sequence-to-Sequence Models"; Correia et al., "Adaptively Sparse Transformers"。
"""
if alpha <= 1.000001:
return F.softmax(z, dim=dim)
# 数值稳定:移除最大值
z = z - z.amax(dim=dim, keepdim=True)
# 二分求阈值 tau 使得 sum p = 1
d = z.size(dim)
tau_lo = z.amin(dim=dim, keepdim=True)
tau_hi = z.amax(dim=dim, keepdim=True)
for _ in range(n_iter):
tau_m = (tau_lo + tau_hi) * 0.5
relu = (z - tau_m).clamp_min(0.0)
p_m = torch.pow(relu, 1.0 / (alpha - 1.0))
s = p_m.sum(dim=dim, keepdim=True)
# 若和>1,说明阈值偏低,需要抬高下界
gt = (s > 1.0)
tau_lo = torch.where(gt, tau_m, tau_lo)
tau_hi = torch.where(gt, tau_hi, tau_m)
relu = (z - tau_hi).clamp_min(0.0)
p = torch.pow(relu, 1.0 / (alpha - 1.0))
# 归一化确保严格和为1
s = p.sum(dim=dim, keepdim=True).clamp_min(1e-12)
p = p / s
return p
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
x: [batch_size, num_stocks, seq_len, d_model]
return: [batch_size, num_stocks, num_stocks]
"""
batch_size, num_stocks, seq_len, d_model = x.shape
# Attention Pooling: 使用注意力权重聚合时间维度
scores = self.time_pooling_weight(x).squeeze(-1) # [B, N, T]
attn_weights = F.softmax(scores, dim=-1).unsqueeze(-1) # [B, N, T, 1]
last_features = (x * attn_weights).sum(dim=2) # [B, N, D]
# 构造两两组合 [B, N, N, 2D]
fi = last_features.unsqueeze(2).expand(-1, num_stocks, num_stocks, -1)
fj = last_features.unsqueeze(1).expand(-1, num_stocks, num_stocks, -1)
pair_feature = torch.cat([fi, fj], dim=-1)
relation_score = self.relation_mlp(pair_feature.reshape(-1, 2 * d_model))
adj_matrix = relation_score.view(batch_size, num_stocks, num_stocks)
# 对称化后进行有界激活,防止边权爆炸
adj_matrix = 0.5 * (adj_matrix + adj_matrix.transpose(1, 2))
adj_matrix = torch.tanh(adj_matrix) * 0.5
self.last_adj_raw_abs_mean = adj_matrix.abs().mean()
self.last_frobenius = torch.mean(adj_matrix * adj_matrix)
eye = torch.eye(num_stocks, device=x.device).unsqueeze(0)
adj_matrix = adj_matrix + eye
# 行归一化,避免数值爆炸
adj_matrix = torch.nan_to_num(adj_matrix, nan=0.0, posinf=0.0, neginf=0.0)
row_sums = adj_matrix.sum(dim=-1, keepdim=True).clamp_min(1e-6)
adj_matrix = adj_matrix / row_sums
self.last_learned_adj = adj_matrix
return adj_matrix
class SpacetimeGNNMAM(nn.Module):
"""时空联合建模的GNN-MAM架构(支持可变股票数)"""
def __init__(
self,
input_dim: int,
num_stocks: int,
seq_len: int = 60,
d_model: int = 128,
num_heads: int = 8,
num_layers: int = 4,
d_ff: int = 512,
num_classes: int = 1,
dropout: float = 0.1,
temp: float = 1.0,
use_adaln_cond: bool = True,
adj_activation: str = "softmax",
entmax_alpha: float = 1.5,
# Multi-CLS 聚合参数(仅 attn):
num_pool_queries: int = 6,
pooling_tau: float = 0.7,
diffusion_head_hidden_dim: Optional[int] = None,
diffusion_head_dropout: Optional[float] = None,
):
super().__init__()
self.d_model = d_model
self.seq_len = seq_len
self.input_dim = input_dim
self.num_stocks = num_stocks
# Multi-CLS 配置(固定 attn)
self.num_pool_queries = int(num_pool_queries)
self.pool_mode = "attn"
self.pool_tau = float(pooling_tau)
# 输入投影
self.input_projection = nn.Linear(input_dim, d_model)
self.input_norm = nn.LayerNorm(d_model)
# Multi-CLS:在时间维前插入 K 个查询
self.cls_tokens = nn.Parameter(torch.randn(1, 1, self.num_pool_queries, d_model))
self.pool_score = nn.Linear(d_model, 1, bias=False)
# 动态图学习器(去除固定 per-stock 嵌入)
self.graph_learner = DynamicGraphLearner(
d_model,
temp=temp,
adj_activation=adj_activation,
entmax_alpha=float(entmax_alpha),
)
# 时空联合Transformer层
# 使用AdaLN条件调制(扩散模型标准设计,让每层感知时间步)
self.transformer_layers = nn.ModuleList([
SpacetimeTransformerBlock(d_model, num_heads, d_ff, seq_len, dropout, use_adaln_cond=use_adaln_cond)
for _ in range(num_layers)
])
# 统一读出层:保留简单线性读出以兼容 forward(X) 用于backtest
# predictor 分支已移除;所有评分输出均通过 diffusion_head
# Diffusion/Score 头:时间步嵌入 + 噪声观测 x_t 注入 + 预测 score 或 ε̂
# 允许通过 config 微调 score 头 MLP 的隐藏维与dropout(不改结构,仅改容量)
self.diffusion_time_mlp = nn.Sequential(
nn.Linear(d_model, d_model),
nn.SiLU(),
nn.Linear(d_model, d_model)
)
self.diffusion_xproj = nn.Linear(1, d_model)
# 显式波动率条件:将 [B,N,1] 的 vol_cond 投影到 D 维,与 t/x_t 条件同维度
self.diffusion_volproj = nn.Linear(1, d_model)
score_hidden = d_model if diffusion_head_hidden_dim is None else int(diffusion_head_hidden_dim)
score_do = dropout if diffusion_head_dropout is None else float(diffusion_head_dropout)
self.diffusion_head = nn.Sequential(
nn.Linear(d_model, score_hidden),
nn.SiLU(),
nn.Dropout(score_do),
nn.Linear(score_hidden, 1)
)
self.dropout = nn.Dropout(dropout)
self.last_entropy_reg = torch.tensor(0.0)
self.last_head_div_reg = torch.tensor(0.0)
self.last_graph_reg = torch.tensor(0.0)
# 可视化:按层缓存CLS注意力(列表长度=层数)
self.debug_temporal_cls_attn = [] # List[Tensor or None], 每个元素形状期望 [B,N,H,T+1]
self.debug_spatial_cls_attn = [] # List[Tensor or None], 每个元素形状期望 [B,H,N,N]
self.save_attention = False
# Final LayerNorm to stabilize Pre-LN stack outputs
self.final_layer_norm = nn.LayerNorm(d_model)
def _encode_features(self, x: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor:
"""
将原始输入编码为每只股票的上下文表征(聚合后的 CLS),返回 [B, N, D]。
同步维护正则统计(熵/头多样性/Per-Asset特征)。
Args:
x: [B, N, T, F] 特征
mask: [B, N, T, F] mask (1=有效, 0=缺失) [mask机制]
"""
batch_size, num_stocks, seq_len, input_dim = x.shape
# Per-Asset Feature Attention 已移除,原始输入直接进入线性投影
# 输入投影
x = self.input_projection(x)
x = self.input_norm(x)
# 添加 K 个 CLS
cls_tokens = self.cls_tokens.expand(batch_size, num_stocks, -1, -1)
x = torch.cat([cls_tokens, x], dim=2)
x = self.dropout(x)
# [mask机制] 生成attention mask: [B,N,T] (1=有效, 0=缺失)
# 从 feature mask [B,N,T,F] 聚合:任一特征有效即认为时间步有效
attn_mask = None
if mask is not None:
# mask: [B,N,T,F] -> attn_mask: [B,N,T+K](K为CLS数量)
attn_mask_seq = (mask.sum(dim=-1) > 0).float() # [B,N,T]
# 为CLS token添加全1 mask(CLS永远有效)
cls_mask = torch.ones(batch_size, num_stocks, self.num_pool_queries, device=x.device, dtype=x.dtype)
attn_mask = torch.cat([cls_mask, attn_mask_seq], dim=2) # [B,N,T+K]
# 宏观跳跃连接 (Macro Skip Connection)
# 将原始输入特征(Input Projection后)保留
x_input_proj = x
# 时空编码
ent_regs = []
div_regs = []
graph_regs = []
if self.save_attention:
self.debug_temporal_cls_attn = []
self.debug_spatial_cls_attn = []
for layer in self.transformer_layers:
if hasattr(layer, 'spacetime_attention'):
# 避免静态类型告警:通过 setattr 注入 Python 属性
setattr(layer.spacetime_attention, 'save_attention', bool(self.save_attention))
setattr(layer.spacetime_attention, 'k_vis', int(self.num_pool_queries))
adj_matrix = self.graph_learner(x)
if getattr(self.graph_learner, "last_adj_raw_abs_mean", None) is not None:
graph_regs.append(self.graph_learner.last_adj_raw_abs_mean)
x = layer(x, adj_matrix, mask=attn_mask)
if getattr(layer, 'last_entropy', None) is not None:
ent_regs.append(layer.last_entropy)
if getattr(layer, 'last_head_div', None) is not None:
div_regs.append(layer.last_head_div)
if self.save_attention:
self.debug_temporal_cls_attn.append(getattr(layer, 'last_temporal_cls_attn', None))
self.debug_spatial_cls_attn.append(getattr(layer, 'last_spatial_cls_attn', None))
# [Macro Skip Connection] 在最终输出前再次叠加原始输入
# 确保深层网络能够直接访问最底层的特征
x = x + x_input_proj
x = self.final_layer_norm(x)
# K 个 CLS 聚合
cls_block = x[:, :, : self.num_pool_queries, :]
if self.pool_mode == "attn":
pool_scores = self.pool_score(cls_block).squeeze(-1)
pool_weights = F.softmax(pool_scores / float(self.pool_tau), dim=-1)
pooled = torch.sum(pool_weights.unsqueeze(-1) * cls_block, dim=2)
else:
pooled = torch.sum(cls_block, dim=2) / float(self.num_pool_queries)
# 正则
if len(ent_regs) > 0:
self.last_entropy_reg = torch.stack(ent_regs).mean()
else:
self.last_entropy_reg = torch.tensor(0.0, device=pooled.device)
if len(div_regs) > 0:
self.last_head_div_reg = torch.stack(div_regs).mean()
else:
self.last_head_div_reg = torch.tensor(0.0, device=pooled.device)
if len(graph_regs) > 0:
self.last_graph_reg = torch.stack(graph_regs).mean()
else:
self.last_graph_reg = pooled.new_tensor(0.0)
return pooled
def _encode_features_cond(self, x: torch.Tensor, x_t: torch.Tensor, t_index: torch.Tensor, vol_cond: Optional[torch.Tensor] = None) -> torch.Tensor:
"""
条件版特征编码:在骨干输入阶段即注入时间步嵌入与加噪目标 x_t 的投影,
使得整个时空编码过程对 (t, x_t) 条件敏感。
返回聚合后的 [B, N, D]。
"""
batch_size, num_stocks, seq_len, input_dim = x.shape
# Per-Asset Feature Attention 已移除
# 输入投影
x = self.input_projection(x) # [B,N,T,D]
x = self.input_norm(x)
# 添加 K 个 CLS
cls_tokens = self.cls_tokens.expand(batch_size, num_stocks, -1, -1) # [B,N,K,D]
x = torch.cat([cls_tokens, x], dim=2) # [B,N,K+T,D]
x = self.dropout(x)
x_input_proj = x
# 条件注入:时间步嵌入 + x_t 投影,广播到所有 token(包含 CLS 与时间位)
# 计算与 pooled 相同维度 D 的嵌入/投影
# 注意:与 forward_score 的条件使用共享同一组投影/MLP,保证一致性
B = batch_size
D = x.shape[-1]
t_emb = self._timestep_embedding(t_index, D) # [B,D]
t_emb = self.diffusion_time_mlp(t_emb) # [B,D]
t_broadcast = t_emb.unsqueeze(1).unsqueeze(2).expand(B, num_stocks, x.shape[2], D) # [B,N,K+T,D]
xproj = self.diffusion_xproj(x_t) # 期望 [B,N,D]
xproj_broadcast = xproj.unsqueeze(2).expand(B, num_stocks, x.shape[2], D) # [B,N,K+T,D]
if (vol_cond is not None) and (vol_cond.numel() > 0):
# 期望 vol_cond 形状为 [B,N,1],与 xproj 同步投影后广播
vol_emb = self.diffusion_volproj(vol_cond) # [B,N,D]
vol_broadcast = vol_emb.unsqueeze(2).expand(B, num_stocks, x.shape[2], D) # [B,N,K+T,D]
x = x + t_broadcast + xproj_broadcast + vol_broadcast
else:
x = x + t_broadcast + xproj_broadcast
# 宏观跳跃连接 (Macro Skip Connection)
# 将原始输入特征(Input Projection后)保留
x_input_proj = x
# 时空编码
ent_regs = []
div_regs = []
if self.save_attention:
self.debug_temporal_cls_attn = []
self.debug_spatial_cls_attn = []
graph_regs = []
for layer in self.transformer_layers:
if hasattr(layer, 'spacetime_attention'):
setattr(layer.spacetime_attention, 'save_attention', bool(self.save_attention))
setattr(layer.spacetime_attention, 'k_vis', int(self.num_pool_queries))
adj_matrix = self.graph_learner(x)
if getattr(self.graph_learner, "last_adj_raw_abs_mean", None) is not None:
graph_regs.append(self.graph_learner.last_adj_raw_abs_mean)
# 条件向量:来自 t_emb 与 xproj 的"股票级"表示,避免混入 token 内容
# 不使用try-except,维度错误时立即报错
Bc, Nc, Dc = xproj.shape
base = xproj + t_emb.unsqueeze(1).expand(Bc, Nc, Dc) # [B,N,D]
if (vol_cond is not None) and (vol_cond.numel() > 0):
vol_emb2 = self.diffusion_volproj(vol_cond) # [B,N,D]
cond_vec = base + vol_emb2
else:
cond_vec = base
x = layer(x, adj_matrix, cond=cond_vec)
if getattr(layer, 'last_entropy', None) is not None:
ent_regs.append(layer.last_entropy)
if getattr(layer, 'last_head_div', None) is not None:
div_regs.append(layer.last_head_div)
if self.save_attention:
self.debug_temporal_cls_attn.append(getattr(layer, 'last_temporal_cls_attn', None))
self.debug_spatial_cls_attn.append(getattr(layer, 'last_spatial_cls_attn', None))
# [Macro Skip Connection] 在最终输出前再次叠加原始输入
# 确保深层网络能够直接访问最底层的特征
x = x + x_input_proj
# [Macro Skip Connection]:条件编码路径也叠加原始输入并做层归一化
x = x + x_input_proj
x = self.final_layer_norm(x)
# K 个 CLS 聚合
cls_block = x[:, :, : self.num_pool_queries, :]
if self.pool_mode == "attn":
pool_scores = self.pool_score(cls_block).squeeze(-1)
pool_weights = F.softmax(pool_scores / float(self.pool_tau), dim=-1)
pooled = torch.sum(pool_weights.unsqueeze(-1) * cls_block, dim=2)
else:
pooled = torch.sum(cls_block, dim=2) / float(self.num_pool_queries)
# 正则项聚合
if len(ent_regs) > 0:
self.last_entropy_reg = torch.stack(ent_regs).mean()
else:
self.last_entropy_reg = torch.tensor(0.0, device=pooled.device)
if len(div_regs) > 0:
self.last_head_div_reg = torch.stack(div_regs).mean()
else:
self.last_head_div_reg = torch.tensor(0.0, device=pooled.device)
if len(graph_regs) > 0:
self.last_graph_reg = torch.stack(graph_regs).mean()
else:
self.last_graph_reg = pooled.new_tensor(0.0)
return pooled
def encode(self, x: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor:
"""返回编码后的特征表示(无预测头)。"""
return self._encode_features(x, mask=mask)
def forward(self, x: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor:
"""
x: [batch_size, num_stocks, seq_len, input_dim]
mask: [batch_size, num_stocks, seq_len, input_dim] (1=有效, 0=缺失)
returns: [batch_size, num_stocks, d_model] 编码后的特征
"""
return self.encode(x, mask=mask)
@staticmethod
def _timestep_embedding(t_index: torch.Tensor, dim: int) -> torch.Tensor:
"""标准正弦余弦时间步嵌入(DDPM 风格)。
t_index: [B] 整数,1..T
return: [B, dim]
"""
device = t_index.device
half = dim // 2
# 归一化到 [0,1]
t = t_index.float().unsqueeze(1)
freqs = torch.exp(
torch.arange(0, half, device=device).float() * (-math.log(10000.0) / max(1, half))
).unsqueeze(0)
angles = t * freqs # [B, half]
emb = torch.cat([torch.sin(angles), torch.cos(angles)], dim=1)
if dim % 2 == 1:
emb = F.pad(emb, (0, 1))
return emb
# 已移除 DDPM:不再暴露 forward_diffusion
def forward_score(
self,
x: torch.Tensor,
x_t: torch.Tensor,
t_index: torch.Tensor,
vol_cond: Optional[torch.Tensor] = None,
vol_cond_as_input: bool = True,
vol_cond_scale_output: bool = False
) -> torch.Tensor:
"""
条件 Score 预测:在特征 x 条件下、给定带噪目标 x_t 与时间步 t,预测 score s = ∇_{x_t} log p_t(x_t | x)。
形状与目标变量一致:[B, N, 1]
Args:
vol_cond_as_input: 是否将 vol_cond 作为条件特征输入网络(软引导,让模型学习)
vol_cond_scale_output: 是否用 vol_cond 直接乘到输出上(硬注入,强制调制)
"""
# 根据开关决定是否传递 vol_cond 到条件编码
vol_cond_input = vol_cond if vol_cond_as_input else None
pooled = self._encode_features_cond(x, x_t, t_index, vol_cond=vol_cond_input) # [B, N, D]
f_theta_raw = self.diffusion_head(pooled) # [B, N, 1]
# 根据开关决定是否用 vol_cond 缩放输出
if vol_cond_scale_output and vol_cond is not None and vol_cond.numel() > 0:
# 硬注入:直接用 vol_cond 缩放 f_theta_raw(分标的独立调制)
score_hat = f_theta_raw * vol_cond # [B, N, 1]
else:
# 不缩放,直接返回原始输出
score_hat = f_theta_raw
return score_hat
# 使用示例(略)
if __name__ == "__main__":
# 简单运行一次
B, N, T, F = 2, 5, 8, 4
x = torch.randn(B, N, T, F)
model = SpacetimeGNNMAM(input_dim=F, num_stocks=N, seq_len=T)
y = model(x)
print(y.shape)