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

使用jacrev计算BertForMaskedLM雅可比矩阵报错求助

报错原因
  • 直接原因:你传入的batch_token_ids是词表索引,默认是torch.int64整数类型,PyTorch 仅支持浮点、复数类型张量开启梯度追踪,整数/布尔类型张量无法设置requires_grad=True。
  • 逻辑错误:离散的token id本身是索引值,没有连续可导的定义,哪怕强行把id转成浮点类型,传入BERT的nn.Embedding层时也会触发类型错误(嵌入层要求输入必须是整数索引),对id求雅可比本身没有数学意义。
修复步骤
  1. 删除所有给整数型input_ids、token_type_ids设置requires_grad=True的代码,这两类张量作为索引保持torch.long整数类型即可。
  2. 调整求导对象:雅可比计算要针对嵌入层输出的连续浮点型词向量,而非离散的输入id。你可以通过one-hot编码乘嵌入权重的方式构造可导的输入嵌入,此时输入是浮点类型,支持梯度计算。
  3. 修正原代码中维度扩展的错误:原代码对输入做了两次unsqueeze(0),会把输入维度变成(1,1,seq_len),不符合BERT输入的形状要求,多余的维度扩展需要删除。
修正后的核心代码
import numpy as np
from transformers import BertTokenizer,BertForMaskedLM
import torch
import torch.nn as nn
# PyTorch 2.0+ 已将functorch合并到torch.func,旧版本可保留原functorch导入
try:
    from torch.func import make_functional, make_functional_with_buffers, vmap, vjp, jvp, jacrev
except ImportError:
    from functorch import make_functional, make_functional_with_buffers, vmap, vjp, jvp, jacrev

device = 'cuda:2'
torch.cuda.empty_cache()

model_name = 'bert-base-chinese'
tokenizer = BertTokenizer.from_pretrained(model_name)
bert_model = BertForMaskedLM.from_pretrained(model_name)
net = bert_model.to(device)
fnet, params, buffers = make_functional_with_buffers(net)
# 取出词嵌入层权重,用于构造可导嵌入
embed_weight = bert_model.bert.embeddings.word_embeddings.weight
vocab_size = embed_weight.shape[0]

def fnet_single(params, x_onehot, y):
    # x_onehot: (seq_len, vocab_size) 浮点型可导张量
    # 手动计算词嵌入,替代Embedding层的整数索引操作
    inputs_embeds = (x_onehot @ embed_weight).unsqueeze(0) # (1, seq_len, hidden_size)
    y = y.unsqueeze(0) # (1, seq_len)
    result = fnet(
        params, buffers,
        input_ids=None,
        token_type_ids=y,
        inputs_embeds=inputs_embeds
    )['logits']
    return result.squeeze(0) # (seq_len, vocab_size)

text = u'大肠杆菌是人和许多动物肠道中最主要的一种细菌'
inputs = tokenizer.encode_plus(text)
segment_ids = inputs['token_type_ids']
token_ids = inputs['input_ids']
length = len(token_ids) - 2
seq_len = len(token_ids)

# 输入id保持整数类型,不开梯度
batch_token_ids = torch.tensor([token_ids] * (2 * length - 1), dtype=torch.long).to(device)
batch_segment_ids = torch.zeros_like(batch_token_ids, dtype=torch.long).to(device)

for i in range(length):
    if i > 0:
        batch_token_ids[2 * i - 1, i] = 103
        batch_token_ids[2 * i - 1, i + 1] = 103
    batch_token_ids[2 * i, i + 1] = 103

# 构造浮点型可导的one-hot输入
batch_x_onehot = torch.zeros(
    (2*length -1, seq_len, vocab_size),
    dtype=torch.float32,
    device=device,
    requires_grad=True
)
# 填充one-hot对应位置
for idx in range(2*length -1):
    for pos in range(seq_len):
        tok_id = batch_token_ids[idx, pos].item()
        batch_x_onehot[idx, pos, tok_id] = 1.0

threshold = 100
word_token_ids = [[token_ids[1]]]
for i in range(1, length):
    x = batch_x_onehot[2 * i]
    y = batch_segment_ids[2*i]
    jacobian1 = jacrev(fnet_single, argnums=1)(params, x, y)
    x = batch_x_onehot[2 * i - 1]
    y = batch_segment_ids[2*i-1]
    jacobian2 = jacrev(fnet_single, argnums=1)(params, x, y)

注意:不要尝试直接把整数id转成浮点型开梯度,这种方式会触发嵌入层的类型错误,完全无法运行。实际运行时需要注意显存占用,BERT词表维度较大,全量雅可比计算很容易触发显存溢出,可根据需求做切片计算。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 12:54:27