为何使用AllenNLP的CRF模块会大幅增加训练与推理耗时?
解决AllenNLP CRF模块训练/推理耗时过高的方案
针对你用AllenNLP CRF导致BERT+BiLSTM+CRF模型速度变慢的问题,给你几个可行的优化方向:
换用简化版CRF实现
AllenNLP的CRF模块附带了不少兼容、调试相关的冗余代码,试试自己实现只保留核心功能的CRF,砍掉不必要的计算:import torch import torch.nn as nn class SimplifiedCRF(nn.Module): def __init__(self, num_tags): super().__init__() self.num_tags = num_tags self.transitions = nn.Parameter(torch.randn(num_tags, num_tags)) self.start_trans = nn.Parameter(torch.randn(num_tags)) self.end_trans = nn.Parameter(torch.randn(num_tags)) def forward(self, emissions, tags, mask=None): if mask is None: mask = torch.ones(emissions.shape[:2], dtype=torch.bool, device=emissions.device) log_likelihood = self._calc_log_likelihood(emissions, tags, mask) return -log_likelihood # 返回负对数似然作为损失 def _calc_log_likelihood(self, emissions, tags, mask): batch_size, seq_len = emissions.shape[:2] # 初始化起始分数 score = self.start_trans[tags[:, 0]] + emissions[:, 0, tags[:, 0]] # 遍历序列计算路径分数 for i in range(1, seq_len): mask_i = mask[:, i] if mask_i.any(): score[mask_i] += self.transitions[tags[mask_i, i-1], tags[mask_i, i]] + emissions[mask_i, i, tags[mask_i, i]] # 加上结束转移分数 seq_ends = mask.sum(dim=1) - 1 last_tags = tags[range(batch_size), seq_ends] score += self.end_trans[last_tags] # 计算归一化项(对数分区函数) log_partition = self._calc_log_partition(emissions, mask) return score - log_partition def _calc_log_partition(self, emissions, mask): batch_size, seq_len = emissions.shape[:2] alpha = self.start_trans + emissions[:, 0] for i in range(1, seq_len): mask_i = mask[:, i] if mask_i.any(): # 用矩阵运算替代循环,加速计算 alpha_mat = alpha[mask_i].unsqueeze(1) + self.transitions + emissions[mask_i, i].unsqueeze(0) alpha[mask_i] = torch.logsumexp(alpha_mat, dim=0) # 加上结束转移后求和 return torch.logsumexp(alpha + self.end_trans, dim=1).sum() def decode(self, emissions, mask=None): if mask is None: mask = torch.ones(emissions.shape[:2], dtype=torch.bool, device=emissions.device) batch_size, seq_len = emissions.shape[:2] delta = self.start_trans + emissions[:, 0] paths = [] # 前向计算最大概率路径 for i in range(1, seq_len): mask_i = mask[:, i] if mask_i.any(): delta_mat = delta[mask_i].unsqueeze(1) + self.transitions + emissions[mask_i, i].unsqueeze(0) max_delta, argmax_delta = delta_mat.max(dim=0) delta[mask_i] = max_delta paths.append(argmax_delta) # 回溯得到最优标签序列 delta += self.end_trans best_tags = [] for idx in range(batch_size): seq_end = mask[idx].sum() - 1 best_tag = delta[idx].argmax().item() best_path = [best_tag] for i in reversed(range(1, seq_end+1)): best_tag = paths[i-1][idx][best_tag].item() best_path.append(best_tag) best_path.reverse() # 补全mask外的位置用0填充 best_path += [0]*(seq_len - len(best_path)) best_tags.append(best_path) return best_tags这个版本只保留了对数似然计算和维特比解码的核心逻辑,没有多余的校验和兼容代码,速度会快很多。
开启PyTorch的性能优化开关
- 打开CuDNN的基准模式,让它自动选最优运算算法:
torch.backends.cudnn.benchmark = True - 用混合精度训练,减少计算量和内存占用,间接提升速度:
scaler = torch.cuda.amp.GradScaler() for batch in dataloader: optimizer.zero_grad() with torch.cuda.amp.autocast(): loss = model(**batch) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()
- 打开CuDNN的基准模式,让它自动选最优运算算法:
替换为轻量第三方CRF库
不想自己写的话,试试torchcrf,这是纯PyTorch实现的轻量CRF,没有AllenNLP的额外功能,速度更快:from torchcrf import CRF class YourModel(nn.Module): def __init__(self, bert_model, hidden_size, num_tags): super().__init__() self.bert = bert_model self.bilstm = nn.LSTM(hidden_size, hidden_size//2, bidirectional=True, batch_first=True) self.classifier = nn.Linear(hidden_size, num_tags) self.crf = CRF(num_tags, batch_first=True) # 替换成torchcrf def forward(self, input_ids, attention_mask, tags=None): bert_out = self.bert(input_ids, attention_mask=attention_mask)[0] lstm_out, _ = self.bilstm(bert_out) emissions = self.classifier(lstm_out) if tags is not None: # 训练时返回损失 return -self.crf(emissions, tags, mask=attention_mask.bool(), reduction='mean') else: # 推理时返回解码结果 return self.crf.decode(emissions, mask=attention_mask.bool())排查具体耗时环节
用PyTorch的性能分析工具定位瓶颈,看看是不是CRF的某个函数拖慢了速度:with torch.profiler.profile( activities=[torch.profiler.ProfilerActivity.CPU, torch.profiler.ProfilerActivity.CUDA], record_shapes=True, with_stack=True ) as prof: for batch in dataloader: optimizer.zero_grad() loss = model(**batch) loss.backward() optimizer.step() break # 只跑一个batch就行 print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=10))看输出里的耗时排名,要是CRF的
_calc_log_partition或者解码函数占比最高,就针对性优化这部分的循环或矩阵运算。
内容的提问来源于stack exchange,提问作者Siddharth Tumre
相关产品推荐
相关产品推荐

