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

如何从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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 14:50:27