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

如何实现仅在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)

关键说明

  1. 状态记录:进入上下文时先检查并记录模型当前是否已处于BetterTransformer模式,避免重复转换引发错误。
  2. 原地转换与还原:使用BetterTransformer.transform()和BetterTransformer.reverse()的原地修改特性,确保模型在上下文结束后被准确还原。
  3. 验证效果:修复后,退出上下文后再次检查model.is_bettertransformer应为False,评估速度会回到原模型的约100it/s,后续trainer.train()的padding问题也会解决。

内容的提问来源于stack exchange,提问作者Arthur Thuy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 19:13:31