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

关于torch.nn.Embedding的max_norm参数使用及官方示例解析

关于nn.Embedding中max_norm参数的示例解析与代码适配

一、文档示例的核心逻辑拆解

先明确示例里的基础设定:

  • n=3是嵌入字典的类别总数,d=5是嵌入向量维度,m=7是后续投影层的输出维度
  • max_norm=True表示:每次通过索引获取嵌入向量时,PyTorch会原地修改embedding.weight中对应索引的行,将其归一化到范数为1(若max_norm设为具体数值,则归一化到该数值)

再逐个解析关键代码:

  • a = embedding.weight.clone() @ W.t():
    克隆整个权重矩阵是为了保留归一化前的原始权重快照。因为后续调用embedding(idx)会原地修改权重,若不克隆,直接用embedding.weight计算的话,后续权重变化会导致a的值被篡改。@ W.t()是模拟把嵌入向量投影到m维空间,对应你理解的torch.nn.Linear层权重运算(W就是Linear层的权重参数)。
  • b = embedding(idx) @ W.t():
    通过索引获取嵌入时,PyTorch自动对选中的向量做归一化,并原地更新embedding.weight中对应行。因此b是归一化后的嵌入向量与W的投影结果。
  • out = (a.unsqueeze(0) + b.unsqueeze(1)):
    这是示例为了演示梯度可导性设计的特殊计算:a.unsqueeze(0)将形状从(3,7)转为(1,3,7),b.unsqueeze(1)将形状从(2,7)转为(2,1,7),相加时触发广播,最终得到(2,3,7)的张量——本质是让每个选中的归一化嵌入,和所有原始嵌入的投影结果做加法,这不是常规业务逻辑,只是文档用来对比原始/修改后权重差异的演示手段。

二、为什么要克隆embedding.weight?

当max_norm非None时,embedding(idx)会原地修改权重矩阵。如果不克隆,先计算embedding.weight @ W.t()再调用embedding(idx),会导致前者的计算结果因为权重被修改而失效,进而破坏梯度计算的正确性。示例里的注释明确标注了这一点:weight must be cloned for this to be differentiable。

三、适配你的代码(无显式W的场景)

你的需求是仅将类别特征嵌入作为网络输入,不需要显式处理W,只需关注以下几点:

  1. 无需克隆整个权重矩阵:文档示例的克隆是特殊场景的需求,你只需要直接获取归一化后的嵌入向量即可。
  2. max_norm的常规用法:
    • 设置max_norm=k(k为正数):每次获取嵌入时,范数超过k的向量会被归一化到范数为k;max_norm=True等价于max_norm=1.0。
    • 原地修改权重是PyTorch默认行为,后续再获取相同索引的嵌入时,会直接使用归一化后的向量。
  3. 极简适配代码示例:
    import torch
    import torch.nn as nn
    
    # 定义嵌入层:100个类别,64维嵌入,归一化到范数2.0
    embedding = nn.Embedding(num_embeddings=100, embedding_dim=64, max_norm=2.0)
    # 你的类别索引输入(比如批量输入3个样本的类别特征)
    idx = torch.tensor([5, 12, 33])
    # 获取归一化后的嵌入向量,直接作为后续网络的输入
    embedded_features = embedding(idx)
    # 后续可将embedded_features传入Linear、Transformer等层继续计算
    
  4. 梯度计算无需额外操作:如果需要对嵌入权重求梯度,直接用embedded_features参与损失计算即可,PyTorch会自动处理原地修改后的权重梯度,不需要克隆。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 03:31:03