使用jacrev计算BertForMaskedLM雅可比矩阵报错求助
报错原因
- 直接原因:你传入的
batch_token_ids是词表索引,默认是torch.int64整数类型,PyTorch 仅支持浮点、复数类型张量开启梯度追踪,整数/布尔类型张量无法设置requires_grad=True。 - 逻辑错误:离散的token id本身是索引值,没有连续可导的定义,哪怕强行把id转成浮点类型,传入BERT的
nn.Embedding层时也会触发类型错误(嵌入层要求输入必须是整数索引),对id求雅可比本身没有数学意义。
修复步骤
- 删除所有给整数型
input_ids、token_type_ids设置requires_grad=True的代码,这两类张量作为索引保持torch.long整数类型即可。 - 调整求导对象:雅可比计算要针对嵌入层输出的连续浮点型词向量,而非离散的输入id。你可以通过one-hot编码乘嵌入权重的方式构造可导的输入嵌入,此时输入是浮点类型,支持梯度计算。
- 修正原代码中维度扩展的错误:原代码对输入做了两次
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
相关产品推荐
相关产品推荐

