Transformer位置编码中maxlen的作用及正确设置方式咨询
关于PyTorch位置编码中maxlen参数的疑问解答
官方实现代码
import math import torch from torch import nn, Tensor class PositionalEncoding(nn.Module): def __init__(self, emb_size: int, dropout: float, maxlen: int = 5000): super(PositionalEncoding, self).__init__() den = torch.exp(- torch.arange(0, emb_size, 2)* math.log(10000) / emb_size) pos = torch.arange(0, maxlen).reshape(maxlen, 1) pos_embedding = torch.zeros((maxlen, emb_size)) pos_embedding[:, 0::2] = torch.sin(pos * den) pos_embedding[:, 1::2] = torch.cos(pos * den) pos_embedding = pos_embedding.unsqueeze(-2) self.dropout = nn.Dropout(dropout) self.register_buffer('pos_embedding', pos_embedding) def forward(self, token_embedding: Tensor): return self.dropout(token_embedding + self.pos_embedding[:token_embedding.size(0), :])
疑问解答
maxlen是固定常量还是需根据batch size或数据长度调整?
maxlen是固定预设的常量,不需要根据batch size调整,只需保证它大于等于你训练/推理时会遇到的最大序列长度即可。位置编码是预先生成好0到maxlen-1位置的编码,forward阶段会根据输入序列的实际长度(token_embedding.size(0))截取对应长度的编码使用,和batch size完全无关——batch size是并行处理的样本数量,位置编码针对的是单个样本内的序列位置,两者维度互不影响。输入维度为[256,64](序列长度256,batch size64)时,maxlen=5000是否合理?需要修改吗?
完全不需要修改。maxlen=5000是预先生成的最大位置编码长度,只要你的输入序列长度(256)小于等于5000,就能直接截取前256个位置编码和token embedding相加。官方设置5000是为了覆盖绝大多数NLP场景的序列长度(常见文本序列很少超过5000个token),你的示例序列长度远小于5000,当前使用方式完全正确,无需调整maxlen的值。
内容的提问来源于stack exchange,提问作者Hope
相关产品推荐
相关产品推荐

