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

如何用for循环高效生成PyTorch的多组nn.Embedding?

批量生成PyTorch Embedding的高效方案

当然可以用循环批量生成,既能简化冗长代码,也方便后续维护(新增/删除列时只需修改配置)。以下是具体实现:

步骤1:定义嵌入维度配置

先把各列对应的嵌入维度整理成字典,替代分散的变量定义:

embedding_config = {
    'A': 3,
    'B': 3,
    'C': 5,
    # 继续添加其他列,比如'D':4, 'E':6...直到'P'
}

步骤2:用循环批量创建Embedding

在Module的__init__方法中,遍历配置字典,通过setattr把每个Embedding绑定到模块实例上(PyTorch会自动识别这些子模块,确保参数被正确追踪):

import torch.nn as nn
import pandas as pd

class Example(nn.Module):
    def __init__(self, df, embedding_config):
        super().__init__()  # 继承nn.Module必须调用父类初始化
        self.embedding_config = embedding_config
        
        # 循环生成所有Embedding
        for col_name, embed_dim in embedding_config.items():
            # 计算该列需要的嵌入数量(类别总数)
            num_embeddings = df[col_name].max() + 1
            # 动态创建并绑定Embedding实例
            setattr(self, f"{col_name}_embedding", nn.Embedding(num_embeddings, embed_dim))

    def forward(self, x):
        # 示例:获取某列的嵌入结果
        a_embed = self.A_embedding(x['A'])
        # 其他列同理...
        return a_embed

步骤3:使用方式

# 假设你的DataFrame已经准备好
df = pd.DataFrame(...)
# 初始化模型
model = Example(df, embedding_config)

注意事项

  • 必须用setattr绑定:如果自行用字典存储Embedding,PyTorch不会将其视为模块的子参数,训练时不会更新权重,也无法正常转移到GPU。
  • 确保类别索引连续:如果DataFrame中列的类别编码不是从0开始的连续整数,需要先做映射(比如用sklearn.preprocessing.LabelEncoder),否则会出现索引越界或浪费嵌入空间的问题。
  • 原代码遗漏了super().__init__():这是继承nn.Module的必要操作,否则模块无法正常工作。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 10:25:22