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

Transformer中d_model含义及翻译场景下参数与张量形状问询

PyTorch Transformer相关问题解答

1. d_model参数的含义

在PyTorch的nn.Transformer中,d_model是整个架构的核心特征维度:

  • 它是词嵌入层的输出维度,也就是每个token被编码后的向量长度;
  • 所有子模块(多头注意力层、前馈神经网络层、残差连接的归一化层)的输入、输出维度都统一为d_model,保证网络数据流维度一致;
  • 必须是nhead(注意力头数)的整数倍,因为多头注意力会把d_model维度的特征拆分成nhead个独立子空间,每个子空间维度为d_model // nhead。

2. 机器翻译场景下三者的关系

  • sequence_length与d_model:两者没有强制数值绑定。sequence_length是单条样本的token序列长度,影响显存占用(序列越长,单步计算张量越大),但只要显存足够,不同长度序列都能适配同一d_model的Transformer。过长序列可能需要优化位置编码,但这和d_model本身无关。
  • vocabulary_size与d_model:词嵌入层是vocabulary_size到d_model的线性映射,关联体现在信息容量上:
    • d_model过小,无法承载词表中不同token的语义差异,会丢失嵌入信息;
    • d_model过大,会造成参数冗余,增加训练成本和过拟合风险;
    • 实践中通常选择d_model为256、512、1024这类业界常用值,同时满足是注意力头数的倍数,兼顾语义表达能力和计算效率。

3. 英→西翻译场景的参数设置与张量形状

参数设置示例

结合给定场景,选择兼顾性能和显存的常用参数组合:

import torch
import torch.nn as nn

# 核心参数配置
d_model = 512  # 8的倍数,适配nhead=8,是机器翻译任务的经典选择
nhead = 8
num_encoder_layers = 6
num_decoder_layers = 6
dim_feedforward = 2048  # 通常设为d_model的4倍
dropout = 0.1
src_vocab_size = 40000
tgt_vocab_size = 30000

# 定义Transformer模型
transformer = nn.Transformer(
    d_model=d_model,
    nhead=nhead,
    num_encoder_layers=num_encoder_layers,
    num_decoder_layers=num_decoder_layers,
    dim_feedforward=dim_feedforward,
    dropout=dropout
)

# 自定义词嵌入与位置编码(Transformer本身不含这两部分)
src_embedding = nn.Embedding(src_vocab_size, d_model)
tgt_embedding = nn.Embedding(tgt_vocab_size, d_model)
# 实际项目中常用正弦/余弦位置编码,此处简化示例
pos_encoder = nn.Parameter(torch.randn(1, 1000, d_model))

张量形状说明

假设训练时batch_size=32:

  • src张量:输入的英文token id序列,PyTorch Transformer要求形状为[src_sequence_length, batch_size],即[400, 32];经过词嵌入和位置编码后,变为[400, 32, 512],作为编码器输入。
  • tgt张量:输入的西班牙语token id序列(训练时取目标序列前n-1个token),形状为[tgt_sequence_length, batch_size],即[300, 32];经过嵌入和位置编码后变为[300, 32, 512],作为解码器输入。
  • out张量:Transformer输出,形状为[tgt_sequence_length, batch_size, tgt_vocab_size],即[300, 32, 30000];可通过softmax转换为每个位置的token概率分布,用于预测目标序列。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 19:38:41