如何实现仅在with上下文内启用BetterTransformer的自定义管理器
问题描述
我编写了一个自定义with上下文管理器,目的是调用trainer.evaluate()时临时将模型转为BetterTransformer模式,但测试发现上下文结束后模型仍处于BetterTransformer模式,导致后续trainer.train()因padding问题训练效果不佳。上下文前后及内部的评估结果显示:上下文结束后的评估速度仍为BetterTransformer模式下的约115it/s(原模型约100it/s)。
自定义上下文管理器代码:
class BetterTransformerContext: """Temporarily replace a model with a BetterTransformer model.""" def __init__(self, model): self.model = model self.original_model = None def __enter__(self): self.original_model = self.model self.model = BetterTransformer.transform(self.model) return self.model def __exit__(self, exc_type, exc_val, exc_tb): self.model = self.original_model # self.model = BetterTransformer.reverse(self.model) # NOTE: same result
测试输出:
========== Without Optimum (-> should be slow) ========== BT before context: False 100%|█████████████████████████████| 204/204 [00:01<00:00, 103.09it/s] 0.3161764705882353 ========== With Optimum (-> should be fast) ========== BT in context: True 100%|█████████████████████████████| 204/204 [00:01<00:00, 116.68it/s] 0.3161764705882353 ========== Without Optimum (-> should be slow) ========== BT after context: True 100%|█████████████████████████████| 204/204 [00:01<00:00, 116.53it/s] 0.3161764705882353
存在最小可复现示例。
解决方案
问题根源
BetterTransformer.transform()是原地修改模型实例,并非返回全新的模型对象。你之前的代码中self.original_model = self.model只是保存了原模型的引用,当执行transform后,该引用指向的对象已经被修改为BetterTransformer模式,因此__exit__中把self.model赋值回self.original_model,本质还是指向同一个已被修改的模型,无法完成还原。
修复后的上下文管理器
class BetterTransformerContext: """Temporarily switch model to BetterTransformer mode and revert on exit.""" def __init__(self, model): self.model = model self._original_bt_state = False def __enter__(self): # 记录进入上下文前的BT模式状态,避免重复转换 self._original_bt_state = hasattr(self.model, "is_bettertransformer") and self.model.is_bettertransformer if not self._original_bt_state: # 原地转换为BetterTransformer模式 self.model = BetterTransformer.transform(self.model) return self.model def __exit__(self, exc_type, exc_val, exc_tb): # 仅当进入上下文前不是BT模式时,才执行还原操作 if not self._original_bt_state: # 原地还原为原模型模式 self.model = BetterTransformer.reverse(self.model)
关键说明
- 状态记录:进入上下文时先检查并记录模型当前是否已处于BetterTransformer模式,避免重复转换引发错误。
- 原地转换与还原:使用
BetterTransformer.transform()和BetterTransformer.reverse()的原地修改特性,确保模型在上下文结束后被准确还原。 - 验证效果:修复后,退出上下文后再次检查
model.is_bettertransformer应为False,评估速度会回到原模型的约100it/s,后续trainer.train()的padding问题也会解决。
内容的提问来源于stack exchange,提问作者Arthur Thuy
相关产品推荐
相关产品推荐

