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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 22:06:15