Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions docs/source/Instruction/Supported-models-and-datasets.md
Original file line number Diff line number Diff line change
Expand Up @@ -498,10 +498,10 @@
|[deepseek-ai/DeepSeek-V3.2-Exp](https://modelscope.cn/models/deepseek-ai/DeepSeek-V3.2-Exp)|deepseek_v32|deepseek_v3_1|-|✔|-|[deepseek-ai/DeepSeek-V3.2-Exp](https://huggingface.co/deepseek-ai/DeepSeek-V3.2-Exp)|
|[deepseek-ai/DeepSeek-V3.2-Exp-Base](https://modelscope.cn/models/deepseek-ai/DeepSeek-V3.2-Exp-Base)|deepseek_v32|deepseek_v3_1|-|✔|-|[deepseek-ai/DeepSeek-V3.2-Exp-Base](https://huggingface.co/deepseek-ai/DeepSeek-V3.2-Exp-Base)|
|[deepseek-ai/DeepSeek-Math-V2](https://modelscope.cn/models/deepseek-ai/DeepSeek-Math-V2)|deepseek_v32|deepseek_v3_1|-|✔|-|[deepseek-ai/DeepSeek-Math-V2](https://huggingface.co/deepseek-ai/DeepSeek-Math-V2)|
|[deepseek-ai/DeepSeek-V4-Flash](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Flash)|deepseek_v4|deepseek_v4|-|✘|-|[deepseek-ai/DeepSeek-V4-Flash](https://huggingface.co/deepseek-ai/DeepSeek-V4-Flash)|
|[deepseek-ai/DeepSeek-V4-Flash-Base](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Flash-Base)|deepseek_v4|deepseek_v4|-|✘|-|[deepseek-ai/DeepSeek-V4-Flash-Base](https://huggingface.co/deepseek-ai/DeepSeek-V4-Flash-Base)|
|[deepseek-ai/DeepSeek-V4-Pro](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Pro)|deepseek_v4|deepseek_v4|-|✘|-|[deepseek-ai/DeepSeek-V4-Pro](https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro)|
|[deepseek-ai/DeepSeek-V4-Pro-Base](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Pro-Base)|deepseek_v4|deepseek_v4|-|✘|-|[deepseek-ai/DeepSeek-V4-Pro-Base](https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro-Base)|
|[deepseek-ai/DeepSeek-V4-Flash](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Flash)|deepseek_v4|deepseek_v4|-|✔|-|[deepseek-ai/DeepSeek-V4-Flash](https://huggingface.co/deepseek-ai/DeepSeek-V4-Flash)|
|[deepseek-ai/DeepSeek-V4-Flash-Base](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Flash-Base)|deepseek_v4|deepseek_v4|-|✔|-|[deepseek-ai/DeepSeek-V4-Flash-Base](https://huggingface.co/deepseek-ai/DeepSeek-V4-Flash-Base)|
|[deepseek-ai/DeepSeek-V4-Pro](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Pro)|deepseek_v4|deepseek_v4|-|✔|-|[deepseek-ai/DeepSeek-V4-Pro](https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro)|
|[deepseek-ai/DeepSeek-V4-Pro-Base](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Pro-Base)|deepseek_v4|deepseek_v4|-|✔|-|[deepseek-ai/DeepSeek-V4-Pro-Base](https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro-Base)|
|[OpenBuddy/openbuddy-llama-65b-v8-bf16](https://modelscope.cn/models/OpenBuddy/openbuddy-llama-65b-v8-bf16)|openbuddy_llama|openbuddy|-|✔|-|[OpenBuddy/openbuddy-llama-65b-v8-bf16](https://huggingface.co/OpenBuddy/openbuddy-llama-65b-v8-bf16)|
|[OpenBuddy/openbuddy-llama2-13b-v8.1-fp16](https://modelscope.cn/models/OpenBuddy/openbuddy-llama2-13b-v8.1-fp16)|openbuddy_llama|openbuddy|-|✔|-|[OpenBuddy/openbuddy-llama2-13b-v8.1-fp16](https://huggingface.co/OpenBuddy/openbuddy-llama2-13b-v8.1-fp16)|
|[OpenBuddy/openbuddy-llama2-70b-v10.1-bf16](https://modelscope.cn/models/OpenBuddy/openbuddy-llama2-70b-v10.1-bf16)|openbuddy_llama|openbuddy|-|✔|-|[OpenBuddy/openbuddy-llama2-70b-v10.1-bf16](https://huggingface.co/OpenBuddy/openbuddy-llama2-70b-v10.1-bf16)|
Expand Down
9 changes: 7 additions & 2 deletions docs/source/Megatron-SWIFT/Command-line-parameters.md
Original file line number Diff line number Diff line change
Expand Up @@ -152,7 +152,7 @@
- 🔥overlap_param_gather: 启用分布式优化器中参数all-gather的重叠(降低DP通信耗时)。默认为False。
- virtual_pipeline_model_parallel_size: 每个流水线并行 rank 的虚拟流水线阶段数量。默认为None。vpp并行,用于减少pp并行的计算空泡,提高GPU利用率,但会略微提高通信量。
- microbatch_group_size_per_vp_stage: 每个虚拟流水线阶段处理的连续微批次数量。默认为None,等于pipeline_model_parallel_size。
- 🔥pipeline_model_parallel_layout: 一个描述自定义流水线(pp/vpp)模型并行布局的字符串。例如:"E|(t|)*3,m|m||L"。其中 E、L、t、m 分别表示嵌入层(embedding)、损失层(loss)、Transformer 解码器层和 MTP 层。阶段之间用 "|" 分隔。重复的阶段或层可以通过乘法表示。逗号仅用于提升可读性(无实际语法作用)。默认值为 None,表示不使用此参数设置布局。
- 🔥pipeline_model_parallel_layout: 一个描述自定义流水线(pp/vpp)模型并行布局的字符串。例如:`"Et*4|(tttt|)*14tmL"`。其中 E、L、t、m 分别表示嵌入层(embedding)、损失层(loss)、Transformer 解码器层和 MTP 层。阶段之间用 "|" 分隔。重复的阶段或层可以通过乘法表示。逗号仅用于提升可读性(无实际语法作用)。默认值为 None,表示不使用此参数设置布局。
- 该参数通常在异构GPU集群上使用。
- 🔥expert_model_parallel_size: 专家并行数,默认为1。
- 🔥expert_tensor_parallel_size: 专家TP并行度。默认值为1。
Expand Down Expand Up @@ -209,9 +209,14 @@
- moe_token_drop_policy: 可选为'probs', 'position'。默认为'probs'。

**DSA参数**
- dsa_indexer_loss_coeff: DSA 索引器 KL 散度损失的系数。设置为 0 可禁用索引器损失。默认为None
- dsa_indexer_loss_coeff: DSA 索引器 KL 散度损失的系数。设置为 0 可禁用索引器损失。默认为`0.`
- dsa_indexer_use_sparse_loss: 是否使用稀疏 DSA 索引器损失。如果为 True,索引器损失将使用 top-k 索引进行计算。默认为False。

**Deepseek-V4**
- csa_dense_mode: 是否对压缩稀疏注意力使用密集模式。若为 `True`,CSA 索引器将被禁用。默认为False。
- use_fused_mhc: 对 mHC 操作使用 cuTile 融合内核。若为 True,将尝试用融合的 cuda.tile(cuTile)自动微分函数替换参考
mHC 模块以在支持的 GPU 上获得更好的性能。需要安装 cuTile;若 cuTile 不可用,该标志将被静默重置为 False 并发出警告。默认为False。
- mhc_recompute_layer_num: 每个 MHC 重计算块的层数。设置后,每 `mhc_recompute_layer_num` 层构成一个重计算块。若为 None,Transformer 块中的所有层共享单个重计算块。默认为None。

**MTP参数**
- mtp_num_layers: 多token预测(MTP)层的数量。MTP将每个位置的预测范围扩展到多个未来token。此MTP实现使用D个顺序模块依次预测D个额外的token。默认为None。(需要"megatron-core>=0.14")
Expand Down
2 changes: 1 addition & 1 deletion docs/source/Megatron-SWIFT/Quick-start.md
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,7 @@ modelscope-registry.us-west-1.cr.aliyuncs.com/modelscope-repo/modelscope:ubuntu2
| transformer-engine | >=2.3 | 2.14.1 | |
| apex | | 0.1 | |
| megatron-core | >=0.15,<0.18 | 0.17.0 | |
| mcore-bridge | >=1.2.0 | | |
| mcore-bridge | >=1.3.0 | | |
| flash-attn | | 2.8.3/3.0.0b1 | |
| transformers | >=4.33 | 4.57.6/5.8.1 | |
| modelscope | >=1.23 | | |
Expand Down
8 changes: 4 additions & 4 deletions docs/source_en/Instruction/Supported-models-and-datasets.md
Original file line number Diff line number Diff line change
Expand Up @@ -499,10 +499,10 @@ The table below introduces the models integrated with ms-swift:
|[deepseek-ai/DeepSeek-V3.2-Exp](https://modelscope.cn/models/deepseek-ai/DeepSeek-V3.2-Exp)|deepseek_v32|deepseek_v3_1|-|&#x2714;|-|[deepseek-ai/DeepSeek-V3.2-Exp](https://huggingface.co/deepseek-ai/DeepSeek-V3.2-Exp)|
|[deepseek-ai/DeepSeek-V3.2-Exp-Base](https://modelscope.cn/models/deepseek-ai/DeepSeek-V3.2-Exp-Base)|deepseek_v32|deepseek_v3_1|-|&#x2714;|-|[deepseek-ai/DeepSeek-V3.2-Exp-Base](https://huggingface.co/deepseek-ai/DeepSeek-V3.2-Exp-Base)|
|[deepseek-ai/DeepSeek-Math-V2](https://modelscope.cn/models/deepseek-ai/DeepSeek-Math-V2)|deepseek_v32|deepseek_v3_1|-|&#x2714;|-|[deepseek-ai/DeepSeek-Math-V2](https://huggingface.co/deepseek-ai/DeepSeek-Math-V2)|
|[deepseek-ai/DeepSeek-V4-Flash](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Flash)|deepseek_v4|deepseek_v4|-|&#x2718;|-|[deepseek-ai/DeepSeek-V4-Flash](https://huggingface.co/deepseek-ai/DeepSeek-V4-Flash)|
|[deepseek-ai/DeepSeek-V4-Flash-Base](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Flash-Base)|deepseek_v4|deepseek_v4|-|&#x2718;|-|[deepseek-ai/DeepSeek-V4-Flash-Base](https://huggingface.co/deepseek-ai/DeepSeek-V4-Flash-Base)|
|[deepseek-ai/DeepSeek-V4-Pro](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Pro)|deepseek_v4|deepseek_v4|-|&#x2718;|-|[deepseek-ai/DeepSeek-V4-Pro](https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro)|
|[deepseek-ai/DeepSeek-V4-Pro-Base](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Pro-Base)|deepseek_v4|deepseek_v4|-|&#x2718;|-|[deepseek-ai/DeepSeek-V4-Pro-Base](https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro-Base)|
|[deepseek-ai/DeepSeek-V4-Flash](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Flash)|deepseek_v4|deepseek_v4|-|&#x2714;|-|[deepseek-ai/DeepSeek-V4-Flash](https://huggingface.co/deepseek-ai/DeepSeek-V4-Flash)|
|[deepseek-ai/DeepSeek-V4-Flash-Base](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Flash-Base)|deepseek_v4|deepseek_v4|-|&#x2714;|-|[deepseek-ai/DeepSeek-V4-Flash-Base](https://huggingface.co/deepseek-ai/DeepSeek-V4-Flash-Base)|
|[deepseek-ai/DeepSeek-V4-Pro](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Pro)|deepseek_v4|deepseek_v4|-|&#x2714;|-|[deepseek-ai/DeepSeek-V4-Pro](https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro)|
|[deepseek-ai/DeepSeek-V4-Pro-Base](https://modelscope.cn/models/deepseek-ai/DeepSeek-V4-Pro-Base)|deepseek_v4|deepseek_v4|-|&#x2714;|-|[deepseek-ai/DeepSeek-V4-Pro-Base](https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro-Base)|
|[OpenBuddy/openbuddy-llama-65b-v8-bf16](https://modelscope.cn/models/OpenBuddy/openbuddy-llama-65b-v8-bf16)|openbuddy_llama|openbuddy|-|&#x2714;|-|[OpenBuddy/openbuddy-llama-65b-v8-bf16](https://huggingface.co/OpenBuddy/openbuddy-llama-65b-v8-bf16)|
|[OpenBuddy/openbuddy-llama2-13b-v8.1-fp16](https://modelscope.cn/models/OpenBuddy/openbuddy-llama2-13b-v8.1-fp16)|openbuddy_llama|openbuddy|-|&#x2714;|-|[OpenBuddy/openbuddy-llama2-13b-v8.1-fp16](https://huggingface.co/OpenBuddy/openbuddy-llama2-13b-v8.1-fp16)|
|[OpenBuddy/openbuddy-llama2-70b-v10.1-bf16](https://modelscope.cn/models/OpenBuddy/openbuddy-llama2-70b-v10.1-bf16)|openbuddy_llama|openbuddy|-|&#x2714;|-|[OpenBuddy/openbuddy-llama2-70b-v10.1-bf16](https://huggingface.co/OpenBuddy/openbuddy-llama2-70b-v10.1-bf16)|
Expand Down
9 changes: 7 additions & 2 deletions docs/source_en/Megatron-SWIFT/Command-line-parameters.md
Original file line number Diff line number Diff line change
Expand Up @@ -158,7 +158,7 @@ For guidance on selecting parallelization strategies, please refer to the [Train
- 🔥overlap_param_gather: Overlap all-gather of parameters in the distributed optimizer (to reduce DP communication time). Default is False.
- virtual_pipeline_model_parallel_size: The number of virtual pipeline stages per pipeline parallel rank. Defaults to None. VPP parallelism is used to reduce computation bubbles in PP parallelism and improve GPU utilization, but will slightly increase communication overhead.
- microbatch_group_size_per_vp_stage: The number of consecutive microbatches processed by each virtual pipeline stage. Defaults to None, which equals `pipeline_model_parallel_size`.
- 🔥pipeline_model_parallel_layout: A string describing a custom pipeline (pp/vpp) model parallel layout. For example: "E|(t|)*3,m|m||L". Here, E, L, t, and m denote the embedding layer, loss layer, Transformer decoder layer, and MTP layer, respectively. Stages are separated by "|". Repeated stages or layers can be expressed using multiplication. Commas are only for cosmetic readability and have no syntactic meaning. The default value is None, indicating that this argument is not used to set the layout.
- 🔥pipeline_model_parallel_layout: A string describing a custom pipeline (pp/vpp) model parallel layout. For example: `"Et*4|(tttt|)*14tmL"`. Here, E, L, t, and m denote the embedding layer, loss layer, Transformer decoder layer, and MTP layer, respectively. Stages are separated by "|". Repeated stages or layers can be expressed using multiplication. Commas are only for cosmetic readability and have no syntactic meaning. The default value is None, indicating that this argument is not used to set the layout.
- This parameter is typically used on heterogeneous GPU clusters.
- 🔥expert_model_parallel_size: The degree of expert parallelism, default is 1.
- 🔥expert_tensor_parallel_size: expert tensor-parallel size. Default is 1.
Expand Down Expand Up @@ -220,9 +220,14 @@ For guidance on selecting parallelization strategies, please refer to the [Train

**DSA Parameters**

- dsa_indexer_loss_coeff: Coefficient for the DSA indexer KL divergence loss. Set to 0 to disable indexer loss. Default is None.
- dsa_indexer_loss_coeff: Coefficient for the DSA indexer KL divergence loss. Set to 0 to disable indexer loss. Default is `0.`.
- dsa_indexer_use_sparse_loss: Whether to use sparse DSA indexer loss. If True, the indexer loss will be computed using the top-k indices. Default is False.

**Deepseek-V4**

- csa_dense_mode: Whether to use dense mode for compressed sparse attention. If `True`, the CSA indexer will be disabled. Defaults to `False`.
- use_fused_mhc: Use cuTile fused kernels for mHC operations. When `True`, attempts to replace the reference mHC modules with fused cuda.tile (cuTile) autograd functions for better performance on supported GPUs. Requires cuTile to be installed; if cuTile is unavailable, the flag is silently reset to `False` and a warning is emitted. Defaults to `False`.
- mhc_recompute_layer_num: Number of layers per MHC recompute block. When set, every `mhc_recompute_layer_num` layers form a recompute block. If `None`, all layers in the transformer block share a single recompute block. Defaults to `None`.

**MTP Parameters**
- mtp_num_layers: Number of Multi-Token Prediction (MTP) layers. MTP extends the prediction scope at each position to multiple future tokens. This MTP implementation uses D sequential modules to sequentially predict D additional tokens. Default is None. (requires "megatron-core>=0.14")
Expand Down
2 changes: 1 addition & 1 deletion docs/source_en/Megatron-SWIFT/Quick-start.md
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,7 @@ Recommended Operating Environment:
| transformer-engine | >=2.3 | 2.14.1 | |
| apex | | 0.1 | |
| megatron-core | >=0.15,<0.18 | 0.17.0 | |
| mcore-bridge | >=1.2.0 | | |
| mcore-bridge | >=1.3.0 | | |
| flash-attn | | 2.8.3/3.0.0b1 | |
| transformers | >=4.33 | 4.57.6/5.8.1 | |
| modelscope | >=1.23 | | |
Expand Down
2 changes: 1 addition & 1 deletion requirements/framework.txt
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@ sortedcontainers>=1.5.9
tensorboard
tiktoken
tqdm
transformers>=4.33,<5.9.0
transformers>=4.33,<5.10.0
transformers_stream_generator
trl>=0.15,<1.0
uvicorn
Expand Down
16 changes: 12 additions & 4 deletions requirements/install_all.sh
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,19 @@
# sh requirements/install_all.sh
# pip install sglang -U
pip install "vllm>=0.5.1" -U
pip install "transformers<5.9" "trl<1.0" peft -U
pip install optimum bitsandbytes "gradio<5.33" -U
pip install "transformers<5.10" "trl<1.0" "peft<0.20" "datasets<4.8.5" -U
pip install optimum bitsandbytes "gradio<5.33" mcore-bridge -U
pip install "ms-swift[all]@git+https://github.com/modelscope/ms-swift.git"
pip install timm "deepspeed<0.19" -U
pip install timm "deepspeed<0.19" ray -U
pip install qwen_vl_utils qwen_omni_utils keye_vl_utils -U
pip install decord librosa icecream soundfile -U
pip install liger_kernel nvitop pre-commit math_verify py-spy wandb swanlab -U
# flash-attn: https://github.com/Dao-AILab/flash-attention/releases
pip install "flash-attn==2.8.3" --no-build-isolation
# megatron
pip install pybind11 git+https://github.com/NVIDIA/TransformerEngine.git@stable --no-build-isolation
pip install git+https://github.com/deepseek-ai/DeepGEMM.git@v2.1.1.post3 --no-build-isolation
pip install -U flash-linear-attention --no-build-isolation
pip install -U git+https://github.com/Dao-AILab/causal-conv1d --no-build-isolation
pip install git+https://github.com/Dao-AILab/fast-hadamard-transform --no-build-isolation
pip install git+https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git@v0.2.0
# apex
2 changes: 1 addition & 1 deletion requirements/megatron.txt
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
mcore-bridge>=1.2.0
mcore-bridge>=1.3.0
megatron-core>=0.15
peft>=0.15
10 changes: 5 additions & 5 deletions swift/megatron/arguments/megatron_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -622,8 +622,12 @@ class MegatronArguments(RLHFMegatronArgumentsMixin, MegatronTunerMixin):
aligner_lr: Optional[float] = None

# dsa
dsa_indexer_loss_coeff: Optional[float] = None
dsa_indexer_loss_coeff: float = 0.
dsa_indexer_use_sparse_loss: bool = False
# deepseek-v4

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

Add a blank line before the # deepseek-v4 comment to maintain consistent section separation and improve readability. The previous section (dsa) ends at line 626, and a blank line would help distinguish the new group of arguments, following the style used for the # other section below.

Suggested change
# deepseek-v4
# deepseek-v4

csa_dense_mode: bool = False
use_fused_mhc: bool = False
mhc_recompute_layer_num: Optional[int] = None

# other
check_model: bool = True
Expand Down Expand Up @@ -662,10 +666,6 @@ def load_args_config(ckpt_dir: Optional[str]) -> Dict[str, Any]:

def _set_default(self):
if self.mlp_padding_free:
if self.sequence_parallel:
require_version(
'mcore-bridge>=1.3.0.dev',
'Please install mcore-bridge>=1.3.0.dev to use mlp_padding_free with sequence parallel.')
if self.context_parallel_size > 1:
require_version(
'mcore-bridge>=1.4.0.dev',
Expand Down
2 changes: 1 addition & 1 deletion swift/megatron/init.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,7 +121,7 @@ def _new_load_inline(*args, **kwargs):


def _patch_mcore_bridge():
require_version('mcore-bridge>=1.2.0', 'please install mcore-bridge via `pip install mcore-bridge -U`')
require_version('mcore-bridge>=1.3.0', 'please install mcore-bridge via `pip install mcore-bridge -U`')
import mcore_bridge
from mcore_bridge import GPTBridge
logger.info(f'mcore_bridge.__version__: {mcore_bridge.__version__}')
Expand Down
3 changes: 0 additions & 3 deletions swift/megatron/trainers/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,10 +23,7 @@ def get_batch_on_this_pp_rank(args, data, vp_stage=None):
batch = to_device(data, get_current_device(), non_blocking=True)
if args.pipeline_model_parallel_size == 1:
return batch
is_pp_first_stage = mpu.is_pipeline_first_stage(ignore_virtual=False, vp_stage=vp_stage)
is_pp_last_stage = mpu.is_pipeline_last_stage(ignore_virtual=False, vp_stage=vp_stage)
if not args.mtp_num_layers and not is_pp_first_stage:
batch['input_ids'] = None
if not is_pp_last_stage:
batch['labels'] = None
batch['loss_scale'] = None
Expand Down
3 changes: 3 additions & 0 deletions swift/megatron/utils/convert_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -252,6 +252,9 @@ def test_convert_precision(args, hf_model, mg_model, template, test_convert_dtyp
for n, m in mg_language_model.named_modules():
if n.endswith('router'):
m.to(mg_dtype)
if getattr(config, 'enable_hyper_connections', False):
for param in mg_language_model.decoder.parameters(recurse=False):
param.data = param.data.cuda()
Comment thread
Jintao-Huang marked this conversation as resolved.
attention_context = (
_patch_attention_fp32(mg_dtype) if args.attention_backend.name in {'flash', 'fused'} else nullcontext())
with torch.inference_mode(), _model_cpu_forward_context(
Expand Down
Loading