You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何修改ViT PyTorch Transformer模型用于回归任务?含股票、房产预测数据集,求代码

如何修改ViT(PyTorch)适配回归任务(股票/房价预测场景)

核心修改要点

ViT原本为图像分类设计,适配回归任务只需调整三个核心部分:

  • 替换输出头:把原模型输出类别数的线性层,改成输出1个(单变量回归,如预测房价)或多个(多变量回归,如预测股票多维度指标)数值的线性层。
  • 更换损失函数:分类用交叉熵,回归改用MSE(均方误差)、MAE(平均绝对误差)或Huber损失(对异常值更鲁棒)。
  • 数据格式适配:将股票/房价这类表格/时间序列数据拆分成ViT可处理的「patch」——把一维特征/时间步序列切分为固定大小的子序列,每个子序列作为Transformer的输入token。

实现代码示例

以下是简化的ViT回归模型实现,可直接用于股票时间序列或房价表格数据:

1. 基础模块定义

import torch
import torch.nn as nn
from einops import rearrange, repeat

# 特征/时间序列转Patch嵌入
class PatchEmbedding(nn.Module):
    def __init__(self, in_features, patch_size, embed_dim):
        super().__init__()
        self.patch_size = patch_size
        # 对每个patch做线性映射得到嵌入向量
        self.proj = nn.Linear(patch_size, embed_dim)
    
    def forward(self, x):
        # x shape: (batch_size, 总特征数/时间步长)
        # 拆分一维序列为多个固定大小的patch
        x = rearrange(x, 'b (n p) -> b n p', p=self.patch_size)
        return self.proj(x)  # 输出shape: (batch_size, patch数量, embed_dim)

# Transformer编码器层
class TransformerEncoderLayer(nn.Module):
    def __init__(self, embed_dim, num_heads, mlp_dim, dropout=0.1):
        super().__init__()
        self.norm1 = nn.LayerNorm(embed_dim)
        self.attn = nn.MultiheadAttention(embed_dim, num_heads, dropout=dropout, batch_first=True)
        self.norm2 = nn.LayerNorm(embed_dim)
        self.mlp = nn.Sequential(
            nn.Linear(embed_dim, mlp_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(mlp_dim, embed_dim),
            nn.Dropout(dropout)
        )
    
    def forward(self, x):
        # 多头注意力+残差连接
        x = x + self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0]
        # MLP+残差连接
        x = x + self.mlp(self.norm2(x))
        return x

2. 完整ViT回归模型

class ViTRegressor(nn.Module):
    def __init__(self, in_features, patch_size, embed_dim, num_heads, num_layers, mlp_dim, num_outputs=1, dropout=0.1):
        super().__init__()
        self.patch_embed = PatchEmbedding(in_features, patch_size, embed_dim)
        num_patches = in_features // patch_size
        
        # CLS Token:用该token的输出做回归预测
        self.cls_token = nn.Parameter(torch.randn(1, 1, embed_dim))
        # 可学习位置编码
        self.pos_embedding = nn.Parameter(torch.randn(1, num_patches + 1, embed_dim))
        self.dropout = nn.Dropout(dropout)
        
        # 堆叠Transformer编码器层
        self.encoder = nn.Sequential(*[
            TransformerEncoderLayer(embed_dim, num_heads, mlp_dim, dropout)
            for _ in range(num_layers)
        ])
        
        self.norm = nn.LayerNorm(embed_dim)
        # 回归头:输出指定数量的预测值
        self.reg_head = nn.Linear(embed_dim, num_outputs)
    
    def forward(self, x):
        batch_size = x.shape[0]
        # 生成patch嵌入
        x = self.patch_embed(x)
        # 为每个样本添加CLS Token
        cls_tokens = repeat(self.cls_token, '1 1 d -> b 1 d', b=batch_size)
        x = torch.cat([cls_tokens, x], dim=1)
        # 添加位置编码
        x += self.pos_embedding
        x = self.dropout(x)
        
        # 经过Transformer编码器
        x = self.encoder(x)
        x = self.norm(x)
        
        # 用CLS Token的输出做回归预测
        cls_output = x[:, 0, :]
        return self.reg_head(cls_output)

3. 模型使用示例

# 假设股票数据集:每个样本包含过去30个交易日的收盘价(单特征时间序列)
in_features = 30
patch_size = 5  # 每个patch包含5个交易日的数据

# 初始化模型
model = ViTRegressor(
    in_features=in_features,
    patch_size=patch_size,
    embed_dim=128,
    num_heads=4,
    num_layers=3,
    mlp_dim=256,
    num_outputs=1  # 预测下一个交易日的收盘价(单变量回归)
)

# 损失函数与优化器
criterion = nn.MSELoss()
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)

# 模拟训练数据
batch_size = 16
x_train = torch.randn(batch_size, in_features)  # 输入:30个时间步的特征
y_train = torch.randn(batch_size, 1)  # 标签:下一个时间步的数值

# 前向传播与训练
outputs = model(x_train)
loss = criterion(outputs, y_train)
loss.backward()
optimizer.step()

实践经验总结

  1. 数据预处理细节
    • 时间序列(股票):用滑动窗口生成样本,比如用过去60天数据预测未来7天价格,窗口大小根据数据频率(日线/小时线)调整。
    • 表格数据(房价):先对离散特征做编码(独热/嵌入),把所有特征拼接成一维向量后拆分patch;也可按特征类别(区位/房屋属性)划分patch。
  2. 模型调优技巧
    • patch大小:时间序列中,patch太小无法捕捉趋势,太大易引入冗余,建议总特征数能被patch大小整除(如30步窗口拆成5或6个patch)。
    • 模型规模:这类数据集不需要大模型,2-4层Transformer、64-256的嵌入维度足够,避免过拟合。
    • 正则化:加入0.1-0.3的dropout、权重衰减,配合早停(Early Stopping)防止过拟合。
  3. 效果对比
    ViT在捕捉长距离依赖上比LSTM/GRU有优势,但需要足够数据量;如果数据集较小,传统模型(XGBoost/LightGBM)可能效果更稳定。

内容的提问来源于stack exchange,提问作者saifhassan

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.29 02:40:42