bf16混合精度下使用nn.TransformerEncoder遇RuntimeError问题求助
问题
在bf16-mixed混合精度下运行模型时触发以下错误:
RuntimeError: expected scalar type Float but found BFloat16
错误出现在使用nn.TransformerEncoder处理Tensor C的环节。
输入包含两种模态数据:
- Tensor H:形状为
(batch_size, 10000),经1D CNN特征提取后传入TransformerEncoder,整个流程在bf16-mixed精度下运行正常。 - Tensor C:形状为
(batch_size, 80),经token化、padding mask生成、nn.Embedding层编码后传入Transformer,此处触发上述类型不匹配错误。
怀疑问题与nn.Embedding在bf16-mixed精度下的行为特性有关,相关代码如下:
# ... other stuff above.. if self.spectrum_type == 'cnmr' and self.cnmr_binary: # Tensor C的处理流程 # Tokenize the binary CNMR data tokens = self._tokenize_cnmr(x) # Create mask for transformer (True values will be ignored) mask = (tokens == 0).bool() # Embed the tokens (this is our feature extraction) x = self.feature_extractor(tokens) # [batch_size, seq_len, d_model] # Apply transformer with mask x = self.transformer(x, src_key_padding_mask=mask) # 此处触发错误 return x, mask else: # Tensor H的处理流程 # Add channel dimension x = x.unsqueeze(1) # [batch_size, 1, sequence_length] # Apply integrated feature extraction (including pooling if configured) x = self.feature_extractor(x) # [batch_size, channels, seq_len] # Reshape for projection x = x.transpose(1, 2) # [batch_size, seq_len, channels] x = self.post_conv_proj(x) # [batch_size, seq_len, d_model] # Add positional encoding if needed if self.use_pos_encoding: x = self.pos_encoding(x) # Apply transformer encoder x = self.transformer(x) # Create empty mask (no padding) mask = torch.zeros(x.shape[0], x.shape[1], dtype=torch.bool, device=x.device) return x, mask # ... other stuff below...
解决方案
1. 显式转换Embedding输出的精度
nn.Embedding默认输出为Float32,在bf16混合精度上下文易出现类型不匹配。在得到Embedding输出后,显式转换为对应精度:
# 转换为bf16匹配混合精度上下文 x = self.feature_extractor(tokens).to(torch.bfloat16) # 若Transformer期望Float32则转成对应类型 # x = self.feature_extractor(tokens).to(torch.float32)
2. 统一Embedding层的权重dtype
初始化或训练前,将self.feature_extractor(即Embedding层)的权重类型与混合精度上下文对齐:
# 初始化时指定 self.feature_extractor = nn.Embedding(num_embeddings, embedding_dim).to(torch.bfloat16) # 或训练前统一转换 self.feature_extractor = self.feature_extractor.to(torch.bfloat16)
3. 确保padding mask的设备一致性
虽然mask为bool类型,仍需保证它与输入张量在同一设备上,避免隐式转换导致的问题:
mask = (tokens == 0).bool().to(x.device)
4. 检查TransformerEncoder的参数dtype
确保TransformerEncoder的所有参数(线性层、层归一化层等)dtype与输入匹配,可统一转换:
self.transformer = self.transformer.to(torch.bfloat16)
5. 完整包裹混合精度上下文
如果使用torch.cuda.amp.autocast,确保Tensor C的整个处理流程都在autocast上下文内执行,避免部分操作在Float32、部分在bf16下运行导致不兼容。
内容的提问来源于stack exchange,提问作者James Arten
相关产品推荐
相关产品推荐

