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

XLA加速场景下Python类方法运行时修改失效问题咨询

失效原因
  1. exec生成的新类无法覆盖原有类引用:你用exec(src, qwen2.__dict__)重新执行类定义,只是在qwen2模块的命名空间里创建了一个新的Qwen2FlashAttention2类对象,但原有的类引用(比如已经被其他模块导入的缓存)不会被更新。而且inspect.getsource读取的是磁盘上的源文件代码,不是内存中修改后的类,所以第二次打印还是原内容。
  2. inspect.getsource的本质限制:这个函数是从磁盘文件读取源码,而非读取内存中类的实时代码,所以哪怕你真的修改了类,它返回的还是原文件的内容。

解决方法

放弃修改源码文本再exec的思路,直接用Python的**猴子补丁(Monkey Patch)**机制修改类的方法,这是动态修改Python类的标准、可靠方式。

方法一:直接重写并替换目标方法

复制原方法的代码,修改需要调整的部分,再替换类的原有方法,确保参数和返回值与原方法完全一致:

def patch_qwen():
    import torch
    import transformers.models.qwen2.modeling_qwen2 as qwen2
    from transformers.models.qwen2.modeling_qwen2 import Qwen2FlashAttention2Output, apply_rotary_pos_emb

    # 保存原方法(可选,用于后续恢复)
    original_forward = qwen2.Qwen2FlashAttention2.forward

    def patched_forward(self, hidden_states, attention_mask=None, position_ids=None, past_key_value=None,
                       output_attentions=False, use_cache=False, cache_position=None, **kwargs):
        bsz, q_len, _ = hidden_states.size()

        kv_seq_len = past_key_value[0].shape[2] if past_key_value is not None else q_len
        if cache_position is not None:
            kv_seq_len = cache_position[-1] + 1
        # 替换目标代码行
        rotary_seq_len = kv_seq_len
        # 原代码:rotary_seq_len = position_ids[:, -1].max().item()

        # 以下是原forward方法的剩余逻辑,直接复制即可
        query_states = self.q_proj(hidden_states)
        key_states = self.k_proj(hidden_states)
        value_states = self.v_proj(hidden_states)

        query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)
        key_states = key_states.view(bsz, kv_seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        value_states = value_states.view(bsz, kv_seq_len, self.num_heads, self.head_dim).transpose(1, 2)

        cos, sin = self.rotary_emb(value_states, seq_len=rotary_seq_len)
        query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)

        if past_key_value is not None:
            key_states = torch.cat([past_key_value[0], key_states], dim=2)
            value_states = torch.cat([past_key_value[1], value_states], dim=2)

        past_key_value = (key_states, value_states) if use_cache else None

        attn_output = self.flash_attention_forward(
            query_states,
            key_states,
            value_states,
            attention_mask,
            q_len=q_len,
            kv_seq_len=kv_seq_len,
            **kwargs,
        )

        attn_output = attn_output.transpose(1, 2).contiguous()
        attn_output = attn_output.view(bsz, q_len, self.hidden_size)
        attn_output = self.o_proj(attn_output)

        attn_weights = None if not output_attentions else None

        return Qwen2FlashAttention2Output(
            last_hidden_state=attn_output,
            past_key_value=past_key_value,
            attentions=attn_weights,
        )

    # 替换类的forward方法
    qwen2.Qwen2FlashAttention2.forward = patched_forward

方法二:动态修改方法源码(更简洁但需注意依赖)

如果不想复制整个方法,可以修改原方法的源码文本,编译后替换类方法:

def patch_qwen():
    import inspect
    import types
    import transformers.models.qwen2.modeling_qwen2 as qwen2

    # 获取原forward方法的源码
    src_lines = inspect.getsource(qwen2.Qwen2FlashAttention2.forward).splitlines()
    replace_str = "        rotary_seq_len = kv_seq_len"
    search_str = "position_ids[:, -1].max().item()"

    # 定位并修改目标行
    for idx, line in enumerate(src_lines):
        if search_str in line:
            src_lines[idx] = replace_str
            break

    # 拼接成完整函数代码
    full_src = '\n'.join(src_lines)
    # 创建包含原模块依赖的命名空间
    namespace = qwen2.__dict__.copy()
    # 执行代码生成新的forward函数
    exec(full_src, namespace)
    patched_forward = namespace['forward']

    # 将函数绑定为类的方法
    qwen2.Qwen2FlashAttention2.forward = types.MethodType(patched_forward, qwen2.Qwen2FlashAttention2)

验证修改效果

修改后可以通过调用方法或打印方法源码(注意inspect.getsource可能仍显示原文件内容,但运行逻辑会是修改后的)来验证:

patch_qwen()
from transformers.models.qwen2.modeling_qwen2 import Qwen2FlashAttention2
# 打印修改后的forward方法
print(inspect.getsource(Qwen2FlashAttention2.forward))
# 或者实例化类调用forward,观察逻辑是否生效

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 06:08:11