Files
ModelHub XC d4e0a1af66 初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD
Source: Original Platform
2026-08-27 23:50:14 +08:00

46 lines
1.4 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import logging
import torch
try:
import deep_ep
from torch_memory_saver import torch_memory_saver
old_init = deep_ep.Buffer.__init__
def new_init(self, *args, **kwargs):
if torch_memory_saver._impl is not None:
torch_memory_saver._impl._binary_wrapper.cdll.tms_set_interesting_region(False)
old_init(self, *args, **kwargs)
torch.cuda.synchronize()
if torch_memory_saver._impl is not None:
torch_memory_saver._impl._binary_wrapper.cdll.tms_set_interesting_region(True)
deep_ep.Buffer.__init__ = new_init
except ImportError:
logging.warning("deep_ep is not installed, some functionalities may be limited.")
try:
from megatron.bridge.models.qwen_vl.modelling_qwen3_vl.text_model import (
Qwen3VLMoETextRotaryEmbedding,
Qwen3VLTextRotaryEmbedding,
)
def patch_rotary_embedding(cls):
_original_forward = cls.forward
def _patched_forward(self, *args, packed_seq_params=None, **kwargs):
return _original_forward(self, *args, **kwargs)
cls.forward = _patched_forward
patch_rotary_embedding(Qwen3VLTextRotaryEmbedding)
patch_rotary_embedding(Qwen3VLMoETextRotaryEmbedding)
except ImportError:
pass
logging.getLogger().setLevel(logging.WARNING)