初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD Source: Original Platform
This commit is contained in:
45
slime/backends/megatron_utils/__init__.py
Normal file
45
slime/backends/megatron_utils/__init__.py
Normal file
@@ -0,0 +1,45 @@
|
||||
# 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)
|
||||
Reference in New Issue
Block a user