如何从model.parameters()列表获取nn.Embedding实例并更新其权重?
问题分析与解决方案
错误根源
你代码中的B = list(A.parameters())得到的是模型可训练参数的列表,每个元素是torch.nn.Parameter(本质是张量)。你的Model类只有node_embed.weight这一个参数,因此B是长度为1的列表。直接对列表B做B[0:10]切片得到的是子列表,而非张量,自然无法执行乘法操作,这就是报错的原因。
要修改self.node_embed的查找表,本质是修改它的weight参数张量,需要先从B中取出这个张量,再对其内容进行修改。
修正后的代码
import torch class Model(torch.nn.Module): def __init__(self, num, dim): super(Model, self).__init__() self.node_embed = torch.nn.Embedding(num, dim) def update_node_embed(B): # 从参数列表中取出node_embed的weight张量(B的第一个元素) embed_weight = B[0] # 对前10行嵌入权重进行原地修改 embed_weight.data[0:10] *= 0.1 # 注意类名首字母大写,实例化模型 A = Model(1000, 16) B = list(A.parameters()) update_node_embed(B) # 验证修改结果 print(A.node_embed.weight[0])
更鲁棒的参数定位方式
如果后续模型添加了其他参数,仅通过索引B[0]获取参数可能会出错。可以通过named_parameters()结合参数名称精准定位目标参数:
def update_node_embed(model): # 直接通过模型实例获取目标参数,无需传递参数列表 for name, param in model.named_parameters(): if name == 'node_embed.weight': param.data[0:10] *= 0.1 break # 调用时直接传入模型实例 update_node_embed(A)
这种方式无需手动处理参数列表,更不易出错。
内容的提问来源于stack exchange,提问作者jason
相关产品推荐
相关产品推荐

