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

PyTorch Embedding层类型报错求助:输入为torch.int64仍提示FloatTensor

问题排查:Embedding层输入类型不匹配报错

问题描述

传入Embedding层的encoder_input数据类型经打印确认是torch.int64,但运行时仍触发如下报错:

RuntimeError: Expected tensor for argument #1 'indices' to have one of the following scalar types: Long, Int; but got torch.FloatTensor instead (while checking arguments for embedding)

训练代码片段

for epoch in range(inital_epoch, config['number_epochs']):
    model.train()
    batch_iterator = tqdm(train_dataloader, desc=f'Processing epoch {epoch:02d}')
    for batch in batch_iterator:
        encoder_input = batch['encoder_input'].to(device) # (batch_size, context_len)
        decoder_input = batch['decoder_input'].to(device) # (batch_size, context_len)
        encoder_mask = batch['encoder_mask'].to(device) # (batch_size, 1, context_len)
        decoder_mask = batch['decoder_mask'].to(device) # (batch_size, context_len, context_len)

        print(f"encoder_input dtype: {encoder_input.dtype}")
        print(f"encoder_input shape: {encoder_input.shape}")

        encoder_output = model.encoder(encoder_input, encoder_mask)

打印输出

encoder_input dtype: torch.int64
encoder_input shape: torch.Size([8, 440])

完整报错栈

Using device cpu
Processing epoch 00:   0%|                                                                                                       | 0/12771 [00:00<?, ?it/s]
encoder_input dtype: torch.int64
encoder_input shape: torch.Size([8, 440])
Processing epoch 00:   0%|                                                                                                       | 0/12771 [00:00<?, ?it/s]
Traceback (most recent call last):
  File "E:\Fahad\College\BE\Project\MT using Transformers\train.py", line 137, in <module>
    train_model(config)
  File "E:\Fahad\College\BE\Project\MT using Transformers\train.py", line 107, in train_model
    encoder_output = model.encoder(encoder_input, encoder_mask) # (batch_size, context_len, embedding_size)
                     ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "E:\Fahad\College\BE\Project\MT using Transformers\model.py", line 246, in encoder
    return self.encoder(src, src_mask)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "E:\Fahad\College\BE\Project\MT using Transformers\model.py", line 240, in encoder
    src = self.input_embedding(src)
          ^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\fahad\AppData\Roaming\Python\Python311\site-packages\torch\nn\modules\module.py", line 1511, in _wrapped_call_impl
    return self._call_impl(*args, **kwargs)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\fahad\AppData\Roaming\Python\Python311\site-packages\torch\nn\modules\module.py", line 1520, in _call_impl
    return forward_call(*args, **kwargs)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "E:\Fahad\College\BE\Project\MT using Transformers\model.py", line 18, in forward
    return self.embedding(x) * math.sqrt(self.embedding_dim)
           ^^^^^^^^^^^^^^^^^
  File "C:\Users\fahad\AppData\Roaming\Python\Python311\site-packages\torch\nn\modules\module.py", line 1511, in _wrapped_call_impl
    return self._call_impl(*args, **kwargs)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\fahad\AppData\Roaming\Python\Python311\site-packages\torch\nn\modules\module.py", line 1520, in _call_impl
    return forward_call(*args, **kwargs)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\fahad\AppData\Roaming\Python\Python311\site-packages\torch\nn\modules\sparse.py", line 163, in forward
    return F.embedding(
           ^^^^^^^^^^^^
  File "C:\Users\fahad\AppData\Roaming\Python\Python311\site-packages\torch\nn\functional.py", line 2237, in embedding
    return torch.embedding(weight, input, padding_idx, scale_grad_by_freq, sparse)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
RuntimeError: Expected tensor for argument #1 'indices' to have one of the following scalar types: Long, Int; but got torch.FloatTensor instead (while checking arguments for embedding)

排查方向

  • 核对模型与输入的设备一致性:确认model是否已经通过model.to(device)转移到和输入相同的设备上。如果模型和输入设备不匹配,PyTorch自动转换时可能会改变张量类型。
  • 检查自定义Embedding层的forward逻辑:查看input_embedding类的forward方法,确认在调用self.embedding(x)之前,有没有对x执行过浮点运算、归一化等会改变类型的操作。
  • 验证DataLoader输出的原始类型:在DataLoader的collate_fn中直接打印数据类型,或者在取出batch后、调用.to(device)之前就打印batch['encoder_input'].dtype,排查是否在设备转移前类型就已经被篡改。
  • 强制转换输入类型做验证:在传入模型前手动执行encoder_input = encoder_input.long(),如果报错消失,说明确实有某个环节偷偷修改了张量类型。
  • 检查Embedding层初始化参数:确认nn.Embedding的padding_idx等参数设置是否合理,虽然这个因素导致类型错误的概率较低,但可以排除。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 07:10:55