Godwind
/

Godwind yiqichen01 commited on
Commit
7986ae3
·
0 Parent(s):

Duplicate from tencent/Hy4-preview

Browse files

Co-authored-by: yiqichen <yiqichen01@users.noreply.huggingface.co>

This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +37 -0
  2. .gitignore +47 -0
  3. LICENSE +210 -0
  4. README.md +241 -0
  5. README_CN.md +232 -0
  6. assets/benchmark-appendix.jpg +3 -0
  7. assets/benchmark.jpg +3 -0
  8. assets/logo-en.png +0 -0
  9. assets/logo-zh.png +0 -0
  10. chat_template.jinja +309 -0
  11. config.json +303 -0
  12. finetune/README.md +385 -0
  13. finetune/README_CN.md +384 -0
  14. finetune/data/example_data.jsonl +8 -0
  15. finetune/deepspeed_support/ds_zero2_no_offload.json +28 -0
  16. finetune/deepspeed_support/ds_zero2_offload.json +34 -0
  17. finetune/deepspeed_support/ds_zero3_no_offload.json +41 -0
  18. finetune/deepspeed_support/ds_zero3_offload.json +41 -0
  19. finetune/deepspeed_support/train.py +603 -0
  20. finetune/deepspeed_support/train.sh +103 -0
  21. finetune/deepspeed_support/train_lora.sh +109 -0
  22. finetune/llama_factory_support/dataset_info.json +16 -0
  23. finetune/llama_factory_support/hy_v4_full_sft.yaml +55 -0
  24. finetune/llama_factory_support/hy_v4_lora_sft.yaml +71 -0
  25. finetune/llama_factory_support/hy_v4_patches.py +768 -0
  26. finetune/llama_factory_support/hy_v4_template.py +50 -0
  27. finetune/llama_factory_support/train_hy_v4.py +99 -0
  28. finetune/llama_factory_support/train_lf.sh +105 -0
  29. finetune/ms_swift_support/fsdp_hy_v4.json +20 -0
  30. finetune/ms_swift_support/hy_v4_full_sft.yaml +63 -0
  31. finetune/ms_swift_support/hy_v4_lora_sft.yaml +65 -0
  32. finetune/ms_swift_support/hy_v4_swift_patches.py +431 -0
  33. finetune/ms_swift_support/sft_train.sh +101 -0
  34. finetune/ms_swift_support/sft_train_lora.sh +104 -0
  35. finetune/requirements.txt +12 -0
  36. generation_config.json +10 -0
  37. model-00001-of-00131.safetensors +3 -0
  38. model-00002-of-00131.safetensors +3 -0
  39. model-00003-of-00131.safetensors +3 -0
  40. model-00004-of-00131.safetensors +3 -0
  41. model-00005-of-00131.safetensors +3 -0
  42. model-00006-of-00131.safetensors +3 -0
  43. model-00007-of-00131.safetensors +3 -0
  44. model-00008-of-00131.safetensors +3 -0
  45. model-00009-of-00131.safetensors +3 -0
  46. model-00010-of-00131.safetensors +3 -0
  47. model-00011-of-00131.safetensors +3 -0
  48. model-00012-of-00131.safetensors +3 -0
  49. model-00013-of-00131.safetensors +3 -0
  50. model-00014-of-00131.safetensors +3 -0
.gitattributes ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ *.7z filter=lfs diff=lfs merge=lfs -text
2
+ *.arrow filter=lfs diff=lfs merge=lfs -text
3
+ *.bin filter=lfs diff=lfs merge=lfs -text
4
+ *.bz2 filter=lfs diff=lfs merge=lfs -text
5
+ *.ckpt filter=lfs diff=lfs merge=lfs -text
6
+ *.ftz filter=lfs diff=lfs merge=lfs -text
7
+ *.gz filter=lfs diff=lfs merge=lfs -text
8
+ *.h5 filter=lfs diff=lfs merge=lfs -text
9
+ *.joblib filter=lfs diff=lfs merge=lfs -text
10
+ *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
+ *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
+ *.model filter=lfs diff=lfs merge=lfs -text
13
+ *.msgpack filter=lfs diff=lfs merge=lfs -text
14
+ *.npy filter=lfs diff=lfs merge=lfs -text
15
+ *.npz filter=lfs diff=lfs merge=lfs -text
16
+ *.onnx filter=lfs diff=lfs merge=lfs -text
17
+ *.ot filter=lfs diff=lfs merge=lfs -text
18
+ *.parquet filter=lfs diff=lfs merge=lfs -text
19
+ *.pb filter=lfs diff=lfs merge=lfs -text
20
+ *.pickle filter=lfs diff=lfs merge=lfs -text
21
+ *.pkl filter=lfs diff=lfs merge=lfs -text
22
+ *.pt filter=lfs diff=lfs merge=lfs -text
23
+ *.pth filter=lfs diff=lfs merge=lfs -text
24
+ *.rar filter=lfs diff=lfs merge=lfs -text
25
+ *.safetensors filter=lfs diff=lfs merge=lfs -text
26
+ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
+ *.tar.* filter=lfs diff=lfs merge=lfs -text
28
+ *.tar filter=lfs diff=lfs merge=lfs -text
29
+ *.tflite filter=lfs diff=lfs merge=lfs -text
30
+ *.tgz filter=lfs diff=lfs merge=lfs -text
31
+ *.wasm filter=lfs diff=lfs merge=lfs -text
32
+ *.xz filter=lfs diff=lfs merge=lfs -text
33
+ *.zip filter=lfs diff=lfs merge=lfs -text
34
+ *.zst filter=lfs diff=lfs merge=lfs -text
35
+ *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ assets/benchmark-appendix.jpg filter=lfs diff=lfs merge=lfs -text
37
+ assets/benchmark.jpg filter=lfs diff=lfs merge=lfs -text
.gitignore ADDED
@@ -0,0 +1,47 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Python
2
+ __pycache__/
3
+ *.py[cod]
4
+ *$py.class
5
+ *.so
6
+ *.egg-info/
7
+ dist/
8
+ build/
9
+ *.egg
10
+
11
+ # Virtual environments
12
+ .venv/
13
+ venv/
14
+ env/
15
+
16
+ # IDE
17
+ .idea/
18
+ .vscode/
19
+ *.swp
20
+ *.swo
21
+
22
+ # OS
23
+ .DS_Store
24
+ Thumbs.db
25
+
26
+ # Model checkpoints (do not commit large model files)
27
+ *.safetensors
28
+ *.bin
29
+ *.pt
30
+ *.pth
31
+ *.ckpt
32
+
33
+ # Training outputs
34
+ output/
35
+ saves/
36
+
37
+ # Generated by the multi-node launch scripts
38
+ env.txt
39
+ hostfile
40
+ pssh.hosts
41
+
42
+ # Logs
43
+ *.log
44
+ logs/
45
+
46
+ # Cache
47
+ .cache/
LICENSE ADDED
@@ -0,0 +1,210 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Tencent is pleased to support the open source community by making Tencent Hy4 preview available.
2
+
3
+ Copyright (C) 2026 Tencent. All rights reserved.
4
+
5
+ Tencent Hy4 preview is licensed under the Apache-2.0.
6
+
7
+
8
+ Terms of the Apache-2.0:
9
+ --------------------------------------------------------------------
10
+ Apache License
11
+ Version 2.0, January 2004
12
+ http://www.apache.org/licenses/
13
+
14
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
15
+
16
+ 1. Definitions.
17
+
18
+ "License" shall mean the terms and conditions for use, reproduction,
19
+ and distribution as defined by Sections 1 through 9 of this document.
20
+
21
+ "Licensor" shall mean the copyright owner or entity authorized by
22
+ the copyright owner that is granting the License.
23
+
24
+ "Legal Entity" shall mean the union of the acting entity and all
25
+ other entities that control, are controlled by, or are under common
26
+ control with that entity. For the purposes of this definition,
27
+ "control" means (i) the power, direct or indirect, to cause the
28
+ direction or management of such entity, whether by contract or
29
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
30
+ outstanding shares, or (iii) beneficial ownership of such entity.
31
+
32
+ "You" (or "Your") shall mean an individual or Legal Entity
33
+ exercising permissions granted by this License.
34
+
35
+ "Source" form shall mean the preferred form for making modifications,
36
+ including but not limited to software source code, documentation
37
+ source, and configuration files.
38
+
39
+ "Object" form shall mean any form resulting from mechanical
40
+ transformation or translation of a Source form, including but
41
+ not limited to compiled object code, generated documentation,
42
+ and conversions to other media types.
43
+
44
+ "Work" shall mean the work of authorship, whether in Source or
45
+ Object form, made available under the License, as indicated by a
46
+ copyright notice that is included in or attached to the work
47
+ (an example is provided in the Appendix below).
48
+
49
+ "Derivative Works" shall mean any work, whether in Source or Object
50
+ form, that is based on (or derived from) the Work and for which the
51
+ editorial revisions, annotations, elaborations, or other modifications
52
+ represent, as a whole, an original work of authorship. For the purposes
53
+ of this License, Derivative Works shall not include works that remain
54
+ separable from, or merely link (or bind by name) to the interfaces of,
55
+ the Work and Derivative Works thereof.
56
+
57
+ "Contribution" shall mean any work of authorship, including
58
+ the original version of the Work and any modifications or additions
59
+ to that Work or Derivative Works thereof, that is intentionally
60
+ submitted to Licensor for inclusion in the Work by the copyright owner
61
+ or by an individual or Legal Entity authorized to submit on behalf of
62
+ the copyright owner. For the purposes of this definition, "submitted"
63
+ means any form of electronic, verbal, or written communication sent
64
+ to the Licensor or its representatives, including but not limited to
65
+ communication on electronic mailing lists, source code control systems,
66
+ and issue tracking systems that are managed by, or on behalf of, the
67
+ Licensor for the purpose of discussing and improving the Work, but
68
+ excluding communication that is conspicuously marked or otherwise
69
+ designated in writing by the copyright owner as "Not a Contribution."
70
+
71
+ "Contributor" shall mean Licensor and any individual or Legal Entity
72
+ on behalf of whom a Contribution has been received by Licensor and
73
+ subsequently incorporated within the Work.
74
+
75
+ 2. Grant of Copyright License. Subject to the terms and conditions of
76
+ this License, each Contributor hereby grants to You a perpetual,
77
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
78
+ copyright license to reproduce, prepare Derivative Works of,
79
+ publicly display, publicly perform, sublicense, and distribute the
80
+ Work and such Derivative Works in Source or Object form.
81
+
82
+ 3. Grant of Patent License. Subject to the terms and conditions of
83
+ this License, each Contributor hereby grants to You a perpetual,
84
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
85
+ (except as stated in this section) patent license to make, have made,
86
+ use, offer to sell, sell, import, and otherwise transfer the Work,
87
+ where such license applies only to those patent claims licensable
88
+ by such Contributor that are necessarily infringed by their
89
+ Contribution(s) alone or by combination of their Contribution(s)
90
+ with the Work to which such Contribution(s) was submitted. If You
91
+ institute patent litigation against any entity (including a
92
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
93
+ or a Contribution incorporated within the Work constitutes direct
94
+ or contributory patent infringement, then any patent licenses
95
+ granted to You under this License for that Work shall terminate
96
+ as of the date such litigation is filed.
97
+
98
+ 4. Redistribution. You may reproduce and distribute copies of the
99
+ Work or Derivative Works thereof in any medium, with or without
100
+ modifications, and in Source or Object form, provided that You
101
+ meet the following conditions:
102
+
103
+ (a) You must give any other recipients of the Work or
104
+ Derivative Works a copy of this License; and
105
+
106
+ (b) You must cause any modified files to carry prominent notices
107
+ stating that You changed the files; and
108
+
109
+ (c) You must retain, in the Source form of any Derivative Works
110
+ that You distribute, all copyright, patent, trademark, and
111
+ attribution notices from the Source form of the Work,
112
+ excluding those notices that do not pertain to any part of
113
+ the Derivative Works; and
114
+
115
+ (d) If the Work includes a "NOTICE" text file as part of its
116
+ distribution, then any Derivative Works that You distribute must
117
+ include a readable copy of the attribution notices contained
118
+ within such NOTICE file, excluding those notices that do not
119
+ pertain to any part of the Derivative Works, in at least one
120
+ of the following places: within a NOTICE text file distributed
121
+ as part of the Derivative Works; within the Source form or
122
+ documentation, if provided along with the Derivative Works; or,
123
+ within a display generated by the Derivative Works, if and
124
+ wherever such third-party notices normally appear. The contents
125
+ of the NOTICE file are for informational purposes only and
126
+ do not modify the License. You may add Your own attribution
127
+ notices within Derivative Works that You distribute, alongside
128
+ or as an addendum to the NOTICE text from the Work, provided
129
+ that such additional attribution notices cannot be construed
130
+ as modifying the License.
131
+
132
+ You may add Your own copyright statement to Your modifications and
133
+ may provide additional or different license terms and conditions
134
+ for use, reproduction, or distribution of Your modifications, or
135
+ for any such Derivative Works as a whole, provided Your use,
136
+ reproduction, and distribution of the Work otherwise complies with
137
+ the conditions stated in this License.
138
+
139
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
140
+ any Contribution intentionally submitted for inclusion in the Work
141
+ by You to the Licensor shall be under the terms and conditions of
142
+ this License, without any additional terms or conditions.
143
+ Notwithstanding the above, nothing herein shall supersede or modify
144
+ the terms of any separate license agreement you may have executed
145
+ with Licensor regarding such Contributions.
146
+
147
+ 6. Trademarks. This License does not grant permission to use the trade
148
+ names, trademarks, service marks, or product names of the Licensor,
149
+ except as required for reasonable and customary use in describing the
150
+ origin of the Work and reproducing the content of the NOTICE file.
151
+
152
+ 7. Disclaimer of Warranty. Unless required by applicable law or
153
+ agreed to in writing, Licensor provides the Work (and each
154
+ Contributor provides its Contributions) on an "AS IS" BASIS,
155
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
156
+ implied, including, without limitation, any warranties or conditions
157
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
158
+ PARTICULAR PURPOSE. You are solely responsible for determining the
159
+ appropriateness of using or redistributing the Work and assume any
160
+ risks associated with Your exercise of permissions under this License.
161
+
162
+ 8. Limitation of Liability. In no event and under no legal theory,
163
+ whether in tort (including negligence), contract, or otherwise,
164
+ unless required by applicable law (such as deliberate and grossly
165
+ negligent acts) or agreed to in writing, shall any Contributor be
166
+ liable to You for damages, including any direct, indirect, special,
167
+ incidental, or consequential damages of any character arising as a
168
+ result of this License or out of the use or inability to use the
169
+ Work (including but not limited to damages for loss of goodwill,
170
+ work stoppage, computer failure or malfunction, or any and all
171
+ other commercial damages or losses), even if such Contributor
172
+ has been advised of the possibility of such damages.
173
+
174
+ 9. Accepting Warranty or Additional Liability. While redistributing
175
+ the Work or Derivative Works thereof, You may choose to offer,
176
+ and charge a fee for, acceptance of support, warranty, indemnity,
177
+ or other liability obligations and/or rights consistent with this
178
+ License. However, in accepting such obligations, You may act only
179
+ on Your own behalf and on Your sole responsibility, not on behalf
180
+ of any other Contributor, and only if You agree to indemnify,
181
+ defend, and hold each Contributor harmless for any liability
182
+ incurred by, or claims asserted against, such Contributor by reason
183
+ of your accepting any such warranty or additional liability.
184
+
185
+ END OF TERMS AND CONDITIONS
186
+
187
+ APPENDIX: How to apply the Apache License to your work.
188
+
189
+ To apply the Apache License to your work, attach the following
190
+ boilerplate notice, with the fields enclosed by brackets "[]"
191
+ replaced with your own identifying information. (Don't include
192
+ the brackets!) The text should be enclosed in the appropriate
193
+ comment syntax for the file format. We also recommend that a
194
+ file or class name and description of purpose be included on the
195
+ same "printed page" as the copyright notice for easier
196
+ identification within third-party archives.
197
+
198
+ Copyright [yyyy] [name of copyright owner]
199
+
200
+ Licensed under the Apache License, Version 2.0 (the "License");
201
+ you may not use this file except in compliance with the License.
202
+ You may obtain a copy of the License at
203
+
204
+ http://www.apache.org/licenses/LICENSE-2.0
205
+
206
+ Unless required by applicable law or agreed to in writing, software
207
+ distributed under the License is distributed on an "AS IS" BASIS,
208
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
209
+ See the License for the specific language governing permissions and
210
+ limitations under the License.
README.md ADDED
@@ -0,0 +1,241 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ library_name: transformers
4
+ pipeline_tag: text-generation
5
+ tags:
6
+ - hunyuan
7
+ - hy4
8
+ - moe
9
+ - text-generation
10
+ ---
11
+ <p align="left">
12
+ <a href="https://huggingface.co/tencent/Hy4-preview/blob/main/README_CN.md">中文</a>&nbsp;|&nbsp;English
13
+ </p>
14
+ <br>
15
+
16
+ <p align="center">
17
+ <img src="assets/logo-en.png" width="400"/> <br>
18
+ </p>
19
+
20
+ <div align="center" style="line-height: 1;">
21
+
22
+
23
+ [![License](https://img.shields.io/badge/License-Apache%202.0-blue)](#license)
24
+ &nbsp;&nbsp;
25
+ [![HuggingFace](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Tencent%20Hy-ffc107?color=ffc107&logoColor=white)](https://huggingface.co/tencent/Hy4-preview)
26
+ &nbsp;&nbsp;
27
+ [![ModelScope](https://img.shields.io/badge/ModelScope-Tencent%20Hy-624aff)](https://modelscope.cn/models/Tencent-Hunyuan/Hy4-preview)
28
+ &nbsp;&nbsp;
29
+ [![cnb.cool](https://img.shields.io/badge/cnb.cool-Tencent%20Hy-blue?logoColor=white)](https://cnb.cool/ai-models/tencent/Hy4-preview)
30
+ &nbsp;&nbsp;
31
+ [![GitCode](https://img.shields.io/badge/GitCode-Tencent%20Hy-red?logoColor=white)](https://ai.gitcode.com/tencent_hunyuan/Hy4-preview)
32
+
33
+ </div>
34
+
35
+ <p align="center">
36
+ 🖥️&nbsp;<a href="https://aistudio.tencent.com/"><b>Official Website</b></a>&nbsp;&nbsp;|&nbsp;&nbsp;
37
+ 💬&nbsp;<a href="https://github.com/Tencent-Hunyuan/Hy4-preview"><b>GitHub</b></a></p>
38
+
39
+ ---
40
+
41
+ ## Table of Contents
42
+
43
+ - [Model Introduction](#model-introduction)
44
+ - [A New Flagship Generation](#a-new-flagship-generation)
45
+ - [Built for Productivity](#built-for-productivity)
46
+ - [Benchmark Appendix](#benchmark-appendix)
47
+ - [Known Limitations](#known-limitations)
48
+ - [News](#news)
49
+ - [Model Links](#model-links)
50
+ - [Quickstart](#quickstart)
51
+ - [Deployment](#deployment)
52
+ - [vLLM](#vllm)
53
+ - [SGLang](#sglang)
54
+ - [Finetuning](#finetuning)
55
+ - [Quantization](#quantization)
56
+ - [License](#license)
57
+ - [Contact Us](#contact-us)
58
+
59
+ ---
60
+
61
+ ## Model Introduction
62
+
63
+ **Hy4 preview** is a new-generation Mixture-of-Experts (MoE) flagship model developed by the Tencent Hy Team. The model comprises 770B total parameters, of which 49B are activated per token. The backbone consists of 78 layers, where the first layer uses a standard dense FFN and the remaining 77 layers replace it with MoE, each containing 256 routed experts and 1 shared expert; every token activates the top-8 routed experts along with the shared expert. In addition to the backbone, 1 native MTP layer (10B total parameters, 0.7B activated) is built in for speculative decoding.
64
+
65
+ On the architecture side, inspired by DeepSeek and GLM, the attention module employs Gated [DeepSeek Sparse Attention](https://arxiv.org/abs/2512.02556) (Gated DSA) with [IndexCache](https://arxiv.org/abs/2603.12201) for cross-layer sparse index reuse. The residual pathway uses [iHC (identity Hyper-Connections)](https://zhuanlan.zhihu.com/p/2010852389670908320) to expand inter-layer information flow.
66
+
67
+ ### Model Specifications
68
+
69
+ > The table below lists backbone parameters only, excluding the MTP layer.
70
+
71
+ | Property | Value |
72
+ |:---|:---|
73
+ | Architecture | Mixture-of-Experts (MoE) |
74
+ | Total Parameters | 770B |
75
+ | Activated Parameters | 49B |
76
+ | Layers | 78 |
77
+ | Hidden Size | 6144 |
78
+ | Attention Type | Gated DSA |
79
+ | Attention Heads | 64 |
80
+ | Query Compression Dimension | 2048 |
81
+ | Key-Value Compression Dimension | 512 |
82
+ | Indexer Heads / Head Dimension | 32 / 128 |
83
+ | Indexer top-k | 2048 |
84
+ | Residual Streams | 4 |
85
+ | Routed Experts | 256 |
86
+ | Shared Experts | 1 |
87
+ | Activated Routed Experts per Token | 8 |
88
+ | MoE Intermediate Size | 2048 |
89
+ | FFN Intermediate Size | 18432 |
90
+ | Context Length | 1M |
91
+ | Vocabulary Size | 120832 |
92
+
93
+ ## A New Flagship Generation
94
+
95
+ We scaled Hy4 preview on three fronts: model size, context length, and training data. Stronger pre-training and a substantially larger post-training run compound into another step change in capability — the largest generation-over-generation gain we've measured, and enough to put Hy4 preview at the open-source frontier.
96
+
97
+ <p align="center">
98
+ <img src="assets/benchmark.jpg" width="100%"/>
99
+ </p>
100
+
101
+ ## Built for Productivity
102
+
103
+ We partnered with top experts inside Tencent — such as software engineers, game developers, finance analysts, and security experts — and built training data around the work they ship. The result is a model that gets meaningfully further on the tasks these teams run every day:
104
+
105
+ **Software engineering**: Better at understanding, planning, debugging, and verifying long-horizon development tasks, with further gains in the visual taste and interaction quality of front-end work.
106
+
107
+ **Office and analysis**: Takes messy context spread across many files and converts it into shareable artifacts — documents, spreadsheets, and presentations — handling data analysis, equations, and financial models with greater precision.
108
+
109
+ **Game development**: Turns a single prompt into a playable prototype and works fluently with game engines, so developers can keep refining complex projects over multiple turns.
110
+
111
+ **Scientific research**: Stronger understanding, reasoning, and problem-solving on hard research questions, with solid progress across AI research, molecular dynamics, condensed matter physics, and pure mathematics.
112
+
113
+ We also continue to co-design Hy4 preview with Tencent products like CodeBuddy and WorkBuddy, so that gains in the model show up in the work people actually do with it. To check that, we ran a blind side-by-side evaluation: 163 internal experts rated model outputs on 203 engineering tasks. Hy4 preview came out slightly ahead of both GLM 5.3 (2.99 vs. 2.92 average, 46.8% wins / 12.8% ties / 40.4% losses) and Kimi K3 (2.99 vs. 2.94, 51.2% wins / 7.9% ties / 40.9% losses).
114
+
115
+ ## Benchmark Appendix
116
+
117
+ <p align="center">
118
+ <img src="assets/benchmark-appendix.jpg" width="100%"/>
119
+ </p>
120
+
121
+ ## Known Limitations
122
+
123
+ This is an early version of Hy4. There is real headroom left in both pre-training and post-training, and we are shipping with known issues — among them, spending longer than necessary reasoning through complex tasks, and a tendency to over-verify its own work. We'll keep iterating quickly on these. As with Hy3 preview, we would rather ship early and hear what breaks — that's what made Hy3 substantially better, and it's how we will get Hy4 right. We will also keep collaborating closely with Tencent's products and in-house experts to push the boundaries of model intelligence while making it more abundant and affordable.
124
+
125
+ ## News
126
+
127
+ * 🔥 We open-source **Hy4 preview** and **Hy4 preview-FP8** model weights on [Hugging Face](https://huggingface.co/tencent/Hy4-preview), [ModelScope](https://modelscope.cn/models/Tencent-Hunyuan/Hy4-preview), [GitCode](https://ai.gitcode.com/tencent_hunyuan/Hy4-preview), and [CNB](https://cnb.cool/ai-models/tencent/Hy4-preview).
128
+
129
+ ## Model Links
130
+
131
+ | Model Name | Description | Hugging Face | ModelScope | GitCode | CNB |
132
+ |:---|:---|:---:|:---:|:---:|:---:|
133
+ | Hy4 preview | Instruct model | 🤗 [Model](https://huggingface.co/tencent/Hy4-preview) | [Model](https://modelscope.cn/models/Tencent-Hunyuan/Hy4-preview) | [Model](https://ai.gitcode.com/tencent_hunyuan/Hy4-preview) | [Model](https://cnb.cool/ai-models/tencent/Hy4-preview) |
134
+ | Hy4 preview-FP8 | FP8 quantized instruct model | 🤗 [Model](https://huggingface.co/tencent/Hy4-preview-FP8) | [Model](https://modelscope.cn/models/Tencent-Hunyuan/Hy4-preview-FP8) | [Model](https://ai.gitcode.com/tencent_hunyuan/Hy4-preview-FP8) | [Model](https://cnb.cool/ai-models/tencent/Hy4-preview-FP8) |
135
+
136
+ ## Quickstart
137
+
138
+ Deploy Hy4 preview with [vLLM](#vllm) or [SGLang](#sglang) first, then call the OpenAI-compatible API:
139
+
140
+ ```python
141
+ from openai import OpenAI
142
+
143
+ client = OpenAI(base_url="http://127.0.0.1:8000/v1", api_key="EMPTY")
144
+
145
+ response = client.chat.completions.create(
146
+ model="hy4-preview",
147
+ messages=[
148
+ {"role": "user", "content": "Hello! Can you briefly introduce yourself?"},
149
+ ],
150
+ temperature=0.9,
151
+ top_p=1.0,
152
+ )
153
+ print(response.choices[0].message.content)
154
+ ```
155
+
156
+ > **Recommended parameters**: `temperature=0.9`, `top_p=1.0`.
157
+ >
158
+ > **Reasoning mode**: Defaults to `"high"` (deep chain-of-thought), which suits complex tasks such as math, coding, and reasoning. For direct responses, pass `extra_body={"chat_template_kwargs": {"reasoning_effort": "no_think"}}`.
159
+
160
+ See the [Deployment](#deployment) section below for how to start the API server.
161
+
162
+ ## Deployment
163
+
164
+ For production serving, we recommend using [vLLM](https://github.com/vllm-project/vllm) or [SGLang](https://docs.sglang.io/).
165
+
166
+ ### vLLM
167
+
168
+ Build vLLM from source:
169
+ ```bash
170
+ uv venv --python 3.12 --seed --managed-python
171
+ source .venv/bin/activate
172
+ git clone https://github.com/vllm-project/vllm.git
173
+ cd vllm
174
+ uv pip install --editable . --torch-backend=auto
175
+ ```
176
+
177
+ Start the vLLM server with MTP enabled:
178
+
179
+ ```bash
180
+ vllm serve tencent/Hy4-preview-FP8 \
181
+ --tensor-parallel-size 8 \
182
+ --speculative-config.method mtp \
183
+ --speculative-config.num_speculative_tokens 3 \
184
+ --attention-backend FLASHMLA_SPARSE \
185
+ --tool-call-parser hy_v4 \
186
+ --reasoning-parser hy_v4 \
187
+ --enable-auto-tool-choice \
188
+ --port 8000 \
189
+ --served-model-name hy4-preview
190
+ ```
191
+
192
+ ### SGLang
193
+
194
+ Build SGLang from source:
195
+ ```bash
196
+ git clone https://github.com/sgl-project/sglang
197
+ cd sglang
198
+ pip3 install pip --upgrade
199
+ pip3 install "transformers>=5.6.0"
200
+ pip3 install -e "python"
201
+ ```
202
+
203
+ Launch SGLang server with MTP enabled:
204
+
205
+ ```bash
206
+ python3 -m sglang.launch_server \
207
+ --model tencent/Hy4-preview-FP8 \
208
+ --tp-size 8 \
209
+ --tool-call-parser hy_v4 \
210
+ --reasoning-parser hy_v4 \
211
+ --speculative-num-steps 2 \
212
+ --speculative-eagle-topk 1 \
213
+ --speculative-num-draft-tokens 3 \
214
+ --speculative-algorithm EAGLE \
215
+ --port 8000 \
216
+ --served-model-name hy4-preview
217
+ ```
218
+
219
+ ## Finetuning
220
+
221
+ Hy4 preview provides a complete model finetuning pipeline. For detailed documentation, please refer to: [Finetuning Guide](https://huggingface.co/tencent/Hy4-preview/blob/main/finetune/README.md)
222
+
223
+ ## Quantization
224
+
225
+ We provide [AngelSlim](https://github.com/tencent/AngelSlim), a more accessible, comprehensive, and efficient toolkit for large model compression. AngelSlim supports a comprehensive suite of compression tools for large-scale multimodal models, including common quantization algorithms, low-bit quantization, and speculative sampling.
226
+
227
+ ## License
228
+
229
+ Hy4 preview is released under the **Apache License 2.0**. See [LICENSE](https://huggingface.co/tencent/Hy4-preview/blob/main/LICENSE) for details.
230
+
231
+ ## Contact Us
232
+
233
+ If you have any questions or suggestions, feel free to reach out to our R&D and product teams via email:
234
+
235
+ 📧 **hunyuan_opensource@tencent.com**
236
+
237
+ ---
238
+
239
+ <p align="center">
240
+ <i>Hy4 preview is developed by the Tencent Hy Team.</i>
241
+ </p>
README_CN.md ADDED
@@ -0,0 +1,232 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ <p align="left">
2
+ <a href="https://huggingface.co/tencent/Hy4-preview">English</a>&nbsp;|&nbsp;中文
3
+ </p>
4
+ <br>
5
+
6
+ <p align="center">
7
+ <img src="assets/logo-zh.png" width="400"/> <br>
8
+ </p>
9
+
10
+ <div align="center" style="line-height: 1;">
11
+
12
+
13
+ [![License](https://img.shields.io/badge/License-Apache%202.0-blue)](#许可证)
14
+ &nbsp;&nbsp;
15
+ [![HuggingFace](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Tencent%20Hy-ffc107?color=ffc107&logoColor=white)](https://huggingface.co/tencent/Hy4-preview)
16
+ &nbsp;&nbsp;
17
+ [![ModelScope](https://img.shields.io/badge/ModelScope-Tencent%20Hy-624aff)](https://modelscope.cn/models/Tencent-Hunyuan/Hy4-preview)
18
+ &nbsp;&nbsp;
19
+ [![cnb.cool](https://img.shields.io/badge/cnb.cool-Tencent%20Hy-blue?logoColor=white)](https://cnb.cool/ai-models/tencent/Hy4-preview)
20
+ &nbsp;&nbsp;
21
+ [![GitCode](https://img.shields.io/badge/GitCode-Tencent%20Hy-red?logoColor=white)](https://ai.gitcode.com/tencent_hunyuan/Hy4-preview)
22
+
23
+ </div>
24
+
25
+ <p align="center">
26
+ 🖥️&nbsp;<a href="https://aistudio.tencent.com/"><b>官方网站</b></a>&nbsp;&nbsp;|&nbsp;&nbsp;
27
+ 💬&nbsp;<a href="https://github.com/Tencent-Hunyuan/Hy4-preview"><b>GitHub</b></a></p>
28
+
29
+ ---
30
+
31
+ ## 目录
32
+
33
+ - [模型介绍](#模型介绍)
34
+ - [新一代旗舰模型](#新一代旗舰模型)
35
+ - [为生产力而生](#为生产力而生)
36
+ - [Benchmark 附录](#benchmark-附录)
37
+ - [已知局限](#已知局限)
38
+ - [新闻](#新闻)
39
+ - [模型链接](#模型链接)
40
+ - [快速开始](#快速开始)
41
+ - [推理和部署](#推理和部署)
42
+ - [vLLM](#vllm)
43
+ - [SGLang](#sglang)
44
+ - [模型微调](#模型微调)
45
+ - [量化工具](#量化工具)
46
+ - [许可证](#许可证)
47
+ - [联系我们](#联系我们)
48
+
49
+ ---
50
+
51
+ ## 模型介绍
52
+
53
+ **Hy4 preview** 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。
54
+
55
+ 在架构层面,受到 DeepSeek 和 GLM 的启发,注意力侧采用 Gated [DeepSeek Sparse Attention](https://arxiv.org/abs/2512.02556)(Gated DSA),并引入 [IndexCache](https://arxiv.org/abs/2603.12201) 跨层复用稀疏索引,残差侧采用 [iHC (identity Hyper-Connections)](https://zhuanlan.zhihu.com/p/2010852389670908320) 扩展层间信息通路。
56
+
57
+ ### 模型规格
58
+
59
+ > 下表仅列出主干网络参数,不含MTP层
60
+
61
+ | 属性 | 值 |
62
+ |:---|:---|
63
+ | 架构 | 混合专家(MoE) |
64
+ | 总参数 | 770B |
65
+ | 激活参数量 | 49B |
66
+ | 层数 | 78 |
67
+ | 隐藏层维度 | 6144 |
68
+ | 注意力类型 | Gated DSA |
69
+ | 注意力头数 | 64 |
70
+ | Query 压缩维度 | 2048 |
71
+ | Key-Value 压缩维度 | 512 |
72
+ | Indexer 头数 / 头维度 | 32 / 128 |
73
+ | Indexer top-k | 2048 |
74
+ | 残差流数 | 4 |
75
+ | 路由专家数 | 256 |
76
+ | 共享专家数 | 1 |
77
+ | 单 token 激活路由专家数 | 8 |
78
+ | MoE中间层维度 | 2048 |
79
+ | FFN中间层维度 | 18432 |
80
+ | 上下文长度 | 1M |
81
+ | 词表大小 | 120832 |
82
+
83
+ ## 新一代旗舰模型
84
+
85
+ Hy4 preview 在模型尺寸、上下文长度、数据规模上都进行了显著的扩展,预训练和后训练的共同进步带来了智能水平的又一次巨大提升,稳居开源模型第一梯队。
86
+
87
+ <p align="center">
88
+ <img src="assets/benchmark.jpg" width="100%"/>
89
+ </p>
90
+
91
+ ## 为生产力而生
92
+
93
+ 通过与腾讯内部软件工程师、游戏开发者、金融分析师、安全专家等各领域顶尖专家的高质量数据共建,Hy4 preview 在各类真实生产力任务上取得显著进步:
94
+
95
+ **软件工程**:增强长程开发任务的理解、规划、调试与验证能力,进一步提升前端开发的视觉审美和交互质量。
96
+
97
+ **办公分析**:显著提升复杂办公环境理解和金融分析能力,着重优化数据分析、跨文件协作,完成从信息处理到文档、表格与演示文稿交付的完整流程。
98
+
99
+ **游戏开发**:增强一句需求直接生成可玩原型的能力,并能熟练使用游戏引擎,开发者可以通过多轮交互持续完善复杂游戏项目。
100
+
101
+ **科学研究**:显著提升复杂科研问题的理解、推理与求解能力,模型在 AI 研发、分子动力学模拟、凝聚态物理、基础数学等各类场景中均有长足进步。
102
+
103
+ 同时,Hy4 preview 持续与 CodeBuddy / WorkBuddy 等产品深度协同,优化生产力场景的真实用户体验。为验证这一点,我们组织了 163 位内部专家基于 203 个真实工程任务进行模型盲测:Hy4 preview(均分 2.99 / 4)略优于 GLM 5.3(均分 2.92 / 4;胜 46.8% / 平 12.8% / 负 40.4%)和 Kimi K3(均分 2.94 / 4;胜 51.2% / 平 7.9% / 负 40.9%)。
104
+
105
+ ## Benchmark 附录
106
+
107
+ <p align="center">
108
+ <img src="assets/benchmark-appendix.jpg" width="100%"/>
109
+ </p>
110
+
111
+ ## 已知局限
112
+
113
+ Hy4 preview 是 Hy4 迭代的一个早期版��,预训练和后训练均仍有较大的提升空间,也有一些已知问题,如复杂任务的长思考和过度自我验证倾向,我们将持续敏捷迭代。如同 Hy3 preview,我们希望通过 Hy4 preview 的尽快发布获得广泛的真实反馈,从而显著提升 Hy4 正式版。同时,我们将坚持发挥与腾讯产品和专家深度合作的独特优势,持续提升生产力的普惠性和上限。
114
+
115
+ ## 新闻
116
+
117
+ * 🔥 我们在 [Hugging Face](https://huggingface.co/tencent/Hy4-preview)、[ModelScope](https://modelscope.cn/models/Tencent-Hunyuan/Hy4-preview)、[GitCode](https://ai.gitcode.com/tencent_hunyuan/Hy4-preview) 和 [CNB](https://cnb.cool/ai-models/tencent/Hy4-preview) 开源了 **Hy4 preview** 和 **Hy4 preview-FP8** 模型权重。
118
+
119
+ ## 模型链接
120
+
121
+ | 模型名 | 简介 | Hugging Face | ModelScope | GitCode | CNB |
122
+ |:---|:---|:---:|:---:|:---:|:---:|
123
+ | Hy4 preview | Instruct 模型 | 🤗 [Model](https://huggingface.co/tencent/Hy4-preview) | [Model](https://modelscope.cn/models/Tencent-Hunyuan/Hy4-preview) | [Model](https://ai.gitcode.com/tencent_hunyuan/Hy4-preview) | [Model](https://cnb.cool/ai-models/tencent/Hy4-preview) |
124
+ | Hy4 preview-FP8 | FP8 量化 Instruct 模型 | 🤗 [Model](https://huggingface.co/tencent/Hy4-preview-FP8) | [Model](https://modelscope.cn/models/Tencent-Hunyuan/Hy4-preview-FP8) | [Model](https://ai.gitcode.com/tencent_hunyuan/Hy4-preview-FP8) | [Model](https://cnb.cool/ai-models/tencent/Hy4-preview-FP8) |
125
+
126
+ ## 快速开始
127
+
128
+ 建议先通过 [vLLM](#vllm) 或 [SGLang](#sglang) 部署服务,然后通过 OpenAI 兼容 API 调用:
129
+
130
+ ```python
131
+ from openai import OpenAI
132
+
133
+ client = OpenAI(base_url="http://127.0.0.1:8000/v1", api_key="EMPTY")
134
+
135
+ response = client.chat.completions.create(
136
+ model="hy4-preview",
137
+ messages=[
138
+ {"role": "user", "content": "你好!请简单介绍一下你自己。"},
139
+ ],
140
+ temperature=0.9,
141
+ top_p=1.0,
142
+ )
143
+ print(response.choices[0].message.content)
144
+ ```
145
+
146
+ > **推荐参数**:`temperature=0.9`,`top_p=1.0`。
147
+ >
148
+ > **推理模式**:默认为 `"high"`(深度思维链),适合数学、编程、推理等复杂任务;日常对话如需直接回复,可传入 `extra_body={"chat_template_kwargs": {"reasoning_effort": "no_think"}}`。
149
+
150
+ 具体部署方式请参考下方[推理和部署](#推理和部署)章节。
151
+
152
+ ## 推理和部署
153
+
154
+ 对于生产环境部署,我们建议使用 [vLLM](https://github.com/vllm-project/vllm) 或 [SGLang](https://docs.sglang.io/)。
155
+
156
+ ### vLLM
157
+
158
+ 从源码构建 vLLM:
159
+
160
+ ```bash
161
+ uv venv --python 3.12 --seed --managed-python
162
+ source .venv/bin/activate
163
+ git clone https://github.com/vllm-project/vllm.git
164
+ cd vllm
165
+ uv pip install --editable . --torch-backend=auto
166
+ ```
167
+
168
+ 启动 vLLM 服务,开启 MTP:
169
+
170
+ ```bash
171
+ vllm serve tencent/Hy4-preview-FP8 \
172
+ --tensor-parallel-size 8 \
173
+ --speculative-config.method mtp \
174
+ --speculative-config.num_speculative_tokens 3 \
175
+ --attention-backend FLASHMLA_SPARSE \
176
+ --tool-call-parser hy_v4 \
177
+ --reasoning-parser hy_v4 \
178
+ --enable-auto-tool-choice \
179
+ --port 8000 \
180
+ --served-model-name hy4-preview
181
+ ```
182
+
183
+ ### SGLang
184
+
185
+ 从源码构建 SGLang:
186
+
187
+ ```bash
188
+ git clone https://github.com/sgl-project/sglang
189
+ cd sglang
190
+ pip3 install pip --upgrade
191
+ pip3 install "transformers>=5.6.0"
192
+ pip3 install -e "python"
193
+ ```
194
+
195
+ 启动 SGLang 服务,开启 MTP:
196
+ ```bash
197
+ python3 -m sglang.launch_server \
198
+ --model tencent/Hy4-preview-FP8 \
199
+ --tp-size 8 \
200
+ --tool-call-parser hy_v4 \
201
+ --reasoning-parser hy_v4 \
202
+ --speculative-num-steps 2 \
203
+ --speculative-eagle-topk 1 \
204
+ --speculative-num-draft-tokens 3 \
205
+ --speculative-algorithm EAGLE \
206
+ --port 8000 \
207
+ --served-model-name hy4-preview
208
+ ```
209
+
210
+ ## 模型微调
211
+
212
+ Hy4 preview 提供了完整的模型微调流程,详细的微调文档请参考:[模型微调指南](https://huggingface.co/tencent/Hy4-preview/blob/main/finetune/README_CN.md)
213
+
214
+ ## 量化工具
215
+
216
+ 我们提供了 [AngelSlim](https://github.com/tencent/AngelSlim)——一套易用、全面、高效的大模型压缩工具包,涵盖常用量化算法、低比特量化和投机采样等能力。
217
+
218
+ ## 许可证
219
+
220
+ Hy4 preview 基于 **Apache 2.0 许可证** 发布。详情请参阅 [LICENSE](https://huggingface.co/tencent/Hy4-preview/blob/main/LICENSE)。
221
+
222
+ ## 联系我们
223
+
224
+ 如有问题或建议,欢迎通过邮件联系我们的研发和产品团队:
225
+
226
+ 📧 **hunyuan_opensource@tencent.com**
227
+
228
+ ---
229
+
230
+ <p align="center">
231
+ <i>Hy4 preview 由腾讯混元团队研发。</i>
232
+ </p>
assets/benchmark-appendix.jpg ADDED

Git LFS Details

  • SHA256: 02c751bf51c52c52cc02094561d7144bbfe3e27084d48d8a405d580da3554e21
  • Pointer size: 132 Bytes
  • Size of remote file: 8.14 MB
assets/benchmark.jpg ADDED

Git LFS Details

  • SHA256: 1380ff67ebed3be637f0288ba7aa956ba4821ddc002100de44996271478c3fd4
  • Pointer size: 133 Bytes
  • Size of remote file: 16.4 MB
assets/logo-en.png ADDED
assets/logo-zh.png ADDED
chat_template.jinja ADDED
@@ -0,0 +1,309 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {#- ----------‑‑‑ special token variables ‑‑‑---------- -#}
2
+ {%- set HYTK = ':opensource' %}
3
+ {%- set hy_start_token = '<|hy_start{}|>'.format(HYTK) %}
4
+ {%- set hy_middle_token = '<|hy_middle{}|>'.format(HYTK) %}
5
+ {%- set hy_end_token = '<|hy_end{}|>'.format(HYTK) %}
6
+ {%- set think_begin_token = '<think{}>'.format(HYTK) %}
7
+ {%- set think_end_token = '</think{}>'.format(HYTK) %}
8
+ {%- set toolcalls_begin_token = '<tool_calls{}>'.format(HYTK) %}
9
+ {%- set toolcalls_end_token = '</tool_calls{}>'.format(HYTK) %}
10
+ {%- set toolcall_begin_token = '<tool_call{}>'.format(HYTK) %}
11
+ {%- set toolcall_end_token = '</tool_call{}>'.format(HYTK) %}
12
+ {%- set argkey_begin_token = '<arg_key{}>'.format(HYTK) %}
13
+ {%- set argkey_end_token = '</arg_key{}>'.format(HYTK) %}
14
+ {%- set argvalue_begin_token = '<arg_value{}>'.format(HYTK) %}
15
+ {%- set argvalue_end_token = '</arg_value{}>'.format(HYTK) %}
16
+ {%- set toolresponse_begin_token = '<tool_response{}>'.format(HYTK) %}
17
+ {%- set toolresponse_end_token = '</tool_response{}>'.format(HYTK) %}
18
+ {%- set reasoning_mode_token = '<|reasoning_mode{}|>'.format(HYTK) %}
19
+
20
+ {#- ----------‑‑‑ hyperparameters variables ‑‑‑---------- -#}
21
+ {%- if not add_generation_prompt is defined %}
22
+ {%- set add_generation_prompt = false %}
23
+ {%- endif %}
24
+ {%- if not preserved_thinking is defined %}
25
+ {%- if not tools %}
26
+ {%- set preserved_thinking = false %}
27
+ {%- else %}
28
+ {%- set preserved_thinking = true %}
29
+ {%- endif %}
30
+ {%- endif %}
31
+
32
+ {%- if not reasoning_effort is defined %}
33
+ {%- set reasoning_effort = 'high' %}
34
+ {%- elif reasoning_effort not in ['high', 'no_think'] %}
35
+ {%- if reasoning_effort is none %}
36
+ {{- raise_exception('reasoning_effort error : None, should be no_think/high') }}
37
+ {%- else %}
38
+ {{- raise_exception('reasoning_effort error : ' + reasoning_effort + ', should be no_think/high') }}
39
+ {%- endif %}
40
+ {%- endif %}
41
+
42
+ {%- if fallback_strategy is defined and fallback_strategy == 'reasoning_toolcall_retry' %}
43
+ {%- set reasoning_effort = 'high' %}
44
+ {%- set add_generation_prompt = false %}
45
+ {%- endif %}
46
+ {%- if not raw_last_assistant is defined %}
47
+ {%- set raw_last_assistant = false %}
48
+ {%- endif %}
49
+
50
+ {%- macro tool_to_json(tool) -%}
51
+ {%- set ns_tool = namespace(first=true) -%}
52
+ {{- '{' -}}
53
+ {%- for k, v in tool.items() -%}
54
+ {%- if k != 'defer_loading' and k != 'strict' -%}
55
+ {%- if not ns_tool.first -%}{{- ', ' -}}{%- endif -%}
56
+ {%- set ns_tool.first = false -%}
57
+ {{- '"' ~ k ~ '": ' ~ (v | tojson(ensure_ascii=False)) -}}
58
+ {%- endif -%}
59
+ {%- endfor -%}
60
+ {{- '}' -}}
61
+ {%- endmacro -%}
62
+
63
+ {%- macro render_content(content) -%}
64
+ {%- if content is string -%}
65
+ {{- content -}}
66
+ {%- elif content is iterable and content is not mapping -%}
67
+ {%- for item in content -%}
68
+ {%- if item is mapping and item.type == 'text' -%}
69
+ {{- item.text -}}
70
+ {%- elif item is string -%}
71
+ {{- item -}}
72
+ {%- else -%}
73
+ {{- item | string -}}
74
+ {%- endif -%}
75
+ {%- endfor -%}
76
+ {%- elif content is none -%}
77
+ {{- '' -}}
78
+ {%- else -%}
79
+ {{- content | string -}}
80
+ {%- endif -%}
81
+ {%- endmacro -%}
82
+
83
+ {%- macro render_tools_prompt() -%}
84
+ {{- '# Tools\n\nYou may call one or more functions to assist with the user query.' -}}
85
+ {{- '\n\nYou are provided with function signatures within <tools></tools> XML tags:' -}}
86
+ {{- '\n<tools>\n' -}}
87
+ {%- set tool_ns = namespace(first=true) -%}
88
+ {%- for tool in tools -%}
89
+ {%- set t = tool['function'] if tool is mapping and 'function' in tool else tool -%}
90
+ {%- if t.defer_loading is not defined or not t.defer_loading -%}
91
+ {%- if not tool_ns.first -%}{{- '\n' -}}{%- endif -%}
92
+ {%- set tool_ns.first = false -%}
93
+ {{- tool_to_json(t) -}}
94
+ {%- endif -%}
95
+ {%- endfor -%}
96
+ {{- '\n</tools>\n\n' -}}
97
+ {{- 'For function call returns, you should first print ' ~ toolcalls_begin_token -}}
98
+ {{- '\nFor each function call, you should return object like:\n' -}}
99
+ {{- toolcall_begin_token ~ '{function-name}' -}}
100
+ {{- argkey_begin_token ~ '{arg-key-1}' ~ argkey_end_token -}}
101
+ {{- argvalue_begin_token ~ '{arg-value-1}' ~ argvalue_end_token -}}
102
+ {{- argkey_begin_token ~ '{arg-key-2}' ~ argkey_end_token -}}
103
+ {{- argvalue_begin_token ~ '{arg-value-2}' ~ argvalue_end_token -}}
104
+ {{- '...' -}}
105
+ {{- toolcall_end_token -}}
106
+ {{- '\nAt the end of function call returns, you should print ' ~ toolcalls_end_token -}}
107
+ {%- endmacro -%}
108
+
109
+ {%- macro render_tool_response(message) -%}
110
+ {%- set content = message['content'] -%}
111
+ {%- if content is string -%}
112
+ {{- toolresponse_begin_token ~ content ~ toolresponse_end_token -}}
113
+ {%- elif content is iterable and content is not mapping and content and content[0] is mapping and content[0].type == 'tool_reference' -%}
114
+ {{- toolresponse_begin_token -}}
115
+ {{- '<tools>\n' -}}
116
+ {%- for tr in content -%}
117
+ {%- for tool in tools -%}
118
+ {%- set t = tool['function'] if tool is mapping and 'function' in tool else tool -%}
119
+ {%- if t.name == tr.name -%}
120
+ {{- tool_to_json(t) ~ '\n' -}}
121
+ {%- endif -%}
122
+ {%- endfor -%}
123
+ {%- endfor -%}
124
+ {{- '</tools>' -}}
125
+ {{- toolresponse_end_token -}}
126
+ {%- elif content is iterable and content is not mapping and content and content[0] is mapping and content[0].output is defined -%}
127
+ {%- for tr in content -%}
128
+ {{- toolresponse_begin_token ~ tr.output ~ toolresponse_end_token -}}
129
+ {%- endfor -%}
130
+ {%- else -%}
131
+ {{- toolresponse_begin_token ~ render_content(content) ~ toolresponse_end_token -}}
132
+ {%- endif -%}
133
+ {%- endmacro -%}
134
+
135
+ {#- consecutive tool messages are clustered into a single tool message whose content is the list of the original tool messages; the result is written to merged_ns.messages (a macro cannot return a list) -#}
136
+ {%- set merged_ns = namespace(messages=[]) %}
137
+ {%- macro merge_tool_responses(messages) -%}
138
+ {%- set merge_ns = namespace(result=[], cluster=[]) -%}
139
+ {%- for message in messages -%}
140
+ {%- if message['role'] == 'tool' -%}
141
+ {%- set merge_ns.cluster = merge_ns.cluster + [message] -%}
142
+ {%- else -%}
143
+ {%- if merge_ns.cluster -%}
144
+ {%- set merge_ns.result = merge_ns.result + [{'role': 'tool', 'content': merge_ns.cluster}] -%}
145
+ {%- set merge_ns.cluster = [] -%}
146
+ {%- endif -%}
147
+ {%- set merge_ns.result = merge_ns.result + [message] -%}
148
+ {%- endif -%}
149
+ {%- endfor -%}
150
+ {%- if merge_ns.cluster -%}
151
+ {%- set merge_ns.result = merge_ns.result + [{'role': 'tool', 'content': merge_ns.cluster}] -%}
152
+ {%- endif -%}
153
+ {%- set merged_ns.messages = merge_ns.result -%}
154
+ {%- endmacro -%}
155
+
156
+ {#- render one tool cluster, reordered to follow the tool_call order of the preceding assistant turn; any missing / ambiguous id falls back to the original tool message order for the whole cluster -#}
157
+ {%- macro render_tool_responses(tool_messages, tool_call_ids) -%}
158
+ {%- set order_ns = namespace(ordered=[], matched=true) -%}
159
+ {%- if tool_call_ids and tool_call_ids | length == tool_messages | length -%}
160
+ {%- for tool_call_id in tool_call_ids -%}
161
+ {%- set hit_ns = namespace(count=0, message=none) -%}
162
+ {%- for tool_message in tool_messages -%}
163
+ {%- set message_id = tool_message['tool_call_id'] if tool_message['tool_call_id'] is defined else tool_message['id'] -%}
164
+ {%- if message_id is defined and message_id == tool_call_id -%}
165
+ {%- set hit_ns.count = hit_ns.count + 1 -%}
166
+ {%- set hit_ns.message = tool_message -%}
167
+ {%- endif -%}
168
+ {%- endfor -%}
169
+ {%- if hit_ns.count == 1 -%}
170
+ {%- set order_ns.ordered = order_ns.ordered + [hit_ns.message] -%}
171
+ {%- else -%}
172
+ {%- set order_ns.matched = false -%}
173
+ {%- endif -%}
174
+ {%- endfor -%}
175
+ {%- else -%}
176
+ {%- set order_ns.matched = false -%}
177
+ {%- endif -%}
178
+ {%- for tool_message in (order_ns.ordered if order_ns.matched else tool_messages) -%}
179
+ {{- render_tool_response(tool_message) -}}
180
+ {%- endfor -%}
181
+ {%- endmacro -%}
182
+
183
+ {%- set _ = merge_tool_responses(messages) %}
184
+ {%- set normed_messages = merged_ns.messages %}
185
+
186
+ {%- set ns = namespace(last_user_index=-1, has_leading_system=false) %}
187
+ {%- if normed_messages and normed_messages[0].role == 'system' %}
188
+ {%- set ns.has_leading_system = true %}
189
+ {%- endif %}
190
+ {%- for message in normed_messages %}
191
+ {%- if message['role'] == 'user' %}
192
+ {%- set ns.last_user_index = loop.index0 %}
193
+ {%- endif %}
194
+ {%- endfor %}
195
+
196
+ {#- no leading system: synthesize one for tools and/or reasoning_mode -#}
197
+ {%- if not ns.has_leading_system %}
198
+ {{- hy_start_token ~ 'system' ~ hy_middle_token -}}
199
+ {%- if tools %}
200
+ {{- render_tools_prompt() -}}
201
+ {%- endif %}
202
+ {{- reasoning_mode_token ~ 'reasoning_effort:' ~ reasoning_effort -}}
203
+ {{- hy_end_token -}}
204
+ {%- endif %}
205
+
206
+ {%- set last_ns = namespace(last_is_assistant=false) %}
207
+ {%- set prev_ns = namespace(tool_call_ids=[]) %}
208
+ {%- for message in normed_messages %}
209
+ {%- if message['role'] == 'system' %}
210
+ {{- hy_start_token ~ 'system' ~ hy_middle_token -}}
211
+ {#- tools / reasoning_mode only attach to the leading system (messages[0]) -#}
212
+ {%- if loop.first %}
213
+ {%- if tools %}
214
+ {{- render_tools_prompt() -}}
215
+ {%- endif %}
216
+ {%- set content = render_content(message['content']) -%}
217
+ {%- if tools and content -%}
218
+ {{- '\n\n' -}}
219
+ {%- endif -%}
220
+ {{- content -}}
221
+ {{- reasoning_mode_token ~ 'reasoning_effort:' ~ reasoning_effort -}}
222
+ {%- else %}
223
+ {{- render_content(message['content']) -}}
224
+ {%- endif %}
225
+ {{- hy_end_token -}}
226
+ {%- elif message['role'] == 'user' %}
227
+ {{- hy_start_token ~ 'user' ~ hy_middle_token -}}
228
+ {{- render_content(message['content']) -}}
229
+ {{- hy_end_token -}}
230
+ {%- elif message['role'] == 'assistant' %}
231
+ {%- set content_body = render_content(message['content']) -%}
232
+ {#- 'reasoning' takes precedence, fall back to 'reasoning_content' when it is empty -#}
233
+ {%- set reasoning_text = '' %}
234
+ {%- if message['reasoning'] is defined and message['reasoning'] is string and message['reasoning'] %}
235
+ {%- set reasoning_text = message['reasoning'] %}
236
+ {%- elif message['reasoning_content'] is defined and message['reasoning_content'] is string and message['reasoning_content'] %}
237
+ {%- set reasoning_text = message['reasoning_content'] %}
238
+ {%- endif %}
239
+ {#- no_think overrides preserved_thinking: history is always emitted with empty think tags -#}
240
+ {%- if reasoning_effort != 'no_think' and (preserved_thinking or loop.index0 > ns.last_user_index) and reasoning_text %}
241
+ {%- set content = think_begin_token ~ reasoning_text ~ think_end_token ~ content_body %}
242
+ {%- else %}
243
+ {%- set content = think_begin_token ~ think_end_token ~ content_body %}
244
+ {%- endif %}
245
+ {{- hy_start_token ~ 'assistant' ~ hy_middle_token -}}
246
+ {%- if message['tool_calls'] is defined and message['tool_calls'] %}
247
+ {#- remember the tool_call ids of this turn to order the following tool responses; an incomplete id set means "do not reorder" -#}
248
+ {%- set ids_ns = namespace(ids=[], complete=true) %}
249
+ {%- for tool in message['tool_calls'] %}
250
+ {%- set tool_call_id = tool['id'] if tool['id'] is defined and tool['id'] else tool['tool_call_id'] %}
251
+ {%- if tool_call_id is defined and tool_call_id %}
252
+ {%- set ids_ns.ids = ids_ns.ids + [tool_call_id] %}
253
+ {%- else %}
254
+ {%- set ids_ns.complete = false %}
255
+ {%- endif %}
256
+ {%- endfor %}
257
+ {%- set prev_ns.tool_call_ids = ids_ns.ids if ids_ns.complete else [] %}
258
+ {{- content -}}
259
+ {{- toolcalls_begin_token -}}
260
+ {%- for tool in message['tool_calls'] -%}
261
+ {%- set func = tool['function'] if tool is mapping and 'function' in tool else tool -%}
262
+ {%- set arguments = func['arguments'] -%}
263
+ {{- toolcall_begin_token ~ func['name'] -}}
264
+ {%- for key, value in arguments.items() -%}
265
+ {{- argkey_begin_token ~ key ~ argkey_end_token -}}
266
+ {%- if value is not string -%}
267
+ {%- set value = value | tojson(ensure_ascii=False) -%}
268
+ {%- endif -%}
269
+ {{- argvalue_begin_token ~ value ~ argvalue_end_token -}}
270
+ {%- endfor -%}
271
+ {{- toolcall_end_token -}}
272
+ {%- endfor -%}
273
+ {{- toolcalls_end_token -}}
274
+ {%- else %}
275
+ {%- if loop.last and raw_last_assistant %}
276
+ {{- content_body -}}
277
+ {%- else %}
278
+ {{- content -}}
279
+ {%- endif %}
280
+ {%- endif %}
281
+ {#- continuation / prefill: last assistant is still open, do not close with hy_end -#}
282
+ {%- if not (loop.last and raw_last_assistant) %}
283
+ {{- hy_end_token -}}
284
+ {%- endif %}
285
+ {%- elif message['role'] == 'tool' %}
286
+ {{- hy_start_token ~ 'tool' ~ hy_middle_token -}}
287
+ {{- render_tool_responses(message['content'], prev_ns.tool_call_ids) -}}
288
+ {{- hy_end_token -}}
289
+ {%- else %}
290
+ {{- hy_start_token ~ message['role'] ~ hy_middle_token -}}
291
+ {{- render_content(message['content']) -}}
292
+ {{- hy_end_token -}}
293
+ {%- endif %}
294
+ {%- if message['role'] != 'assistant' or not (message['tool_calls'] is defined and message['tool_calls']) %}
295
+ {%- set prev_ns.tool_call_ids = [] %}
296
+ {%- endif %}
297
+ {%- if loop.last and message['role'] == 'assistant' %}
298
+ {%- set last_ns.last_is_assistant = true %}
299
+ {%- endif %}
300
+ {%- endfor %}
301
+ {%- if add_generation_prompt %}
302
+ {%- if not last_ns.last_is_assistant %}
303
+ {%- if reasoning_effort == 'no_think' %}
304
+ {{- hy_start_token ~ 'assistant' ~ hy_middle_token ~ think_begin_token ~ think_end_token -}}
305
+ {%- else %}
306
+ {{- hy_start_token ~ 'assistant' ~ hy_middle_token ~ think_begin_token -}}
307
+ {%- endif %}
308
+ {%- endif %}
309
+ {%- endif %}
config.json ADDED
@@ -0,0 +1,303 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "HYV4ForCausalLM"
4
+ ],
5
+ "attention_bias": false,
6
+ "attention_dropout": 0.0,
7
+ "bitwise_backward_align": false,
8
+ "bos_token_id": 120000,
9
+ "dtype": "bfloat16",
10
+ "enable_ihc": true,
11
+ "enable_lm_head_fp32": true,
12
+ "eos_token_id": 120025,
13
+ "gated_mla": true,
14
+ "gating_type": "elementwise",
15
+ "hc_eps": 1e-06,
16
+ "hc_magnitude": 2.0,
17
+ "hc_mult": 4,
18
+ "head_dim": 64,
19
+ "hidden_act": "silu",
20
+ "hidden_size": 6144,
21
+ "index_head_dim": 128,
22
+ "index_n_heads": 32,
23
+ "index_topk": 2048,
24
+ "indexer_types": [
25
+ "full",
26
+ "full",
27
+ "shared",
28
+ "shared",
29
+ "shared",
30
+ "full",
31
+ "shared",
32
+ "shared",
33
+ "shared",
34
+ "full",
35
+ "shared",
36
+ "shared",
37
+ "shared",
38
+ "full",
39
+ "shared",
40
+ "shared",
41
+ "shared",
42
+ "full",
43
+ "shared",
44
+ "shared",
45
+ "shared",
46
+ "full",
47
+ "shared",
48
+ "shared",
49
+ "shared",
50
+ "full",
51
+ "shared",
52
+ "shared",
53
+ "shared",
54
+ "full",
55
+ "shared",
56
+ "shared",
57
+ "shared",
58
+ "full",
59
+ "shared",
60
+ "shared",
61
+ "shared",
62
+ "full",
63
+ "shared",
64
+ "shared",
65
+ "shared",
66
+ "full",
67
+ "shared",
68
+ "shared",
69
+ "shared",
70
+ "full",
71
+ "shared",
72
+ "shared",
73
+ "shared",
74
+ "full",
75
+ "shared",
76
+ "shared",
77
+ "shared",
78
+ "full",
79
+ "shared",
80
+ "shared",
81
+ "shared",
82
+ "full",
83
+ "shared",
84
+ "shared",
85
+ "shared",
86
+ "full",
87
+ "shared",
88
+ "shared",
89
+ "shared",
90
+ "full",
91
+ "shared",
92
+ "shared",
93
+ "shared",
94
+ "full",
95
+ "shared",
96
+ "shared",
97
+ "shared",
98
+ "full",
99
+ "shared",
100
+ "shared",
101
+ "shared",
102
+ "full"
103
+ ],
104
+ "initializer_range": 0.006,
105
+ "intermediate_size": 18432,
106
+ "kv_lora_rank": 512,
107
+ "layer_types": [
108
+ "deepseek_sparse_attention",
109
+ "deepseek_sparse_attention",
110
+ "deepseek_sparse_attention",
111
+ "deepseek_sparse_attention",
112
+ "deepseek_sparse_attention",
113
+ "deepseek_sparse_attention",
114
+ "deepseek_sparse_attention",
115
+ "deepseek_sparse_attention",
116
+ "deepseek_sparse_attention",
117
+ "deepseek_sparse_attention",
118
+ "deepseek_sparse_attention",
119
+ "deepseek_sparse_attention",
120
+ "deepseek_sparse_attention",
121
+ "deepseek_sparse_attention",
122
+ "deepseek_sparse_attention",
123
+ "deepseek_sparse_attention",
124
+ "deepseek_sparse_attention",
125
+ "deepseek_sparse_attention",
126
+ "deepseek_sparse_attention",
127
+ "deepseek_sparse_attention",
128
+ "deepseek_sparse_attention",
129
+ "deepseek_sparse_attention",
130
+ "deepseek_sparse_attention",
131
+ "deepseek_sparse_attention",
132
+ "deepseek_sparse_attention",
133
+ "deepseek_sparse_attention",
134
+ "deepseek_sparse_attention",
135
+ "deepseek_sparse_attention",
136
+ "deepseek_sparse_attention",
137
+ "deepseek_sparse_attention",
138
+ "deepseek_sparse_attention",
139
+ "deepseek_sparse_attention",
140
+ "deepseek_sparse_attention",
141
+ "deepseek_sparse_attention",
142
+ "deepseek_sparse_attention",
143
+ "deepseek_sparse_attention",
144
+ "deepseek_sparse_attention",
145
+ "deepseek_sparse_attention",
146
+ "deepseek_sparse_attention",
147
+ "deepseek_sparse_attention",
148
+ "deepseek_sparse_attention",
149
+ "deepseek_sparse_attention",
150
+ "deepseek_sparse_attention",
151
+ "deepseek_sparse_attention",
152
+ "deepseek_sparse_attention",
153
+ "deepseek_sparse_attention",
154
+ "deepseek_sparse_attention",
155
+ "deepseek_sparse_attention",
156
+ "deepseek_sparse_attention",
157
+ "deepseek_sparse_attention",
158
+ "deepseek_sparse_attention",
159
+ "deepseek_sparse_attention",
160
+ "deepseek_sparse_attention",
161
+ "deepseek_sparse_attention",
162
+ "deepseek_sparse_attention",
163
+ "deepseek_sparse_attention",
164
+ "deepseek_sparse_attention",
165
+ "deepseek_sparse_attention",
166
+ "deepseek_sparse_attention",
167
+ "deepseek_sparse_attention",
168
+ "deepseek_sparse_attention",
169
+ "deepseek_sparse_attention",
170
+ "deepseek_sparse_attention",
171
+ "deepseek_sparse_attention",
172
+ "deepseek_sparse_attention",
173
+ "deepseek_sparse_attention",
174
+ "deepseek_sparse_attention",
175
+ "deepseek_sparse_attention",
176
+ "deepseek_sparse_attention",
177
+ "deepseek_sparse_attention",
178
+ "deepseek_sparse_attention",
179
+ "deepseek_sparse_attention",
180
+ "deepseek_sparse_attention",
181
+ "deepseek_sparse_attention",
182
+ "deepseek_sparse_attention",
183
+ "deepseek_sparse_attention",
184
+ "deepseek_sparse_attention",
185
+ "deepseek_sparse_attention"
186
+ ],
187
+ "learnable_sink": true,
188
+ "learnable_sink_init": 0.0,
189
+ "max_position_embeddings": 1048576,
190
+ "mlp_layer_types": [
191
+ "dense",
192
+ "sparse",
193
+ "sparse",
194
+ "sparse",
195
+ "sparse",
196
+ "sparse",
197
+ "sparse",
198
+ "sparse",
199
+ "sparse",
200
+ "sparse",
201
+ "sparse",
202
+ "sparse",
203
+ "sparse",
204
+ "sparse",
205
+ "sparse",
206
+ "sparse",
207
+ "sparse",
208
+ "sparse",
209
+ "sparse",
210
+ "sparse",
211
+ "sparse",
212
+ "sparse",
213
+ "sparse",
214
+ "sparse",
215
+ "sparse",
216
+ "sparse",
217
+ "sparse",
218
+ "sparse",
219
+ "sparse",
220
+ "sparse",
221
+ "sparse",
222
+ "sparse",
223
+ "sparse",
224
+ "sparse",
225
+ "sparse",
226
+ "sparse",
227
+ "sparse",
228
+ "sparse",
229
+ "sparse",
230
+ "sparse",
231
+ "sparse",
232
+ "sparse",
233
+ "sparse",
234
+ "sparse",
235
+ "sparse",
236
+ "sparse",
237
+ "sparse",
238
+ "sparse",
239
+ "sparse",
240
+ "sparse",
241
+ "sparse",
242
+ "sparse",
243
+ "sparse",
244
+ "sparse",
245
+ "sparse",
246
+ "sparse",
247
+ "sparse",
248
+ "sparse",
249
+ "sparse",
250
+ "sparse",
251
+ "sparse",
252
+ "sparse",
253
+ "sparse",
254
+ "sparse",
255
+ "sparse",
256
+ "sparse",
257
+ "sparse",
258
+ "sparse",
259
+ "sparse",
260
+ "sparse",
261
+ "sparse",
262
+ "sparse",
263
+ "sparse",
264
+ "sparse",
265
+ "sparse",
266
+ "sparse",
267
+ "sparse",
268
+ "sparse"
269
+ ],
270
+ "model_type": "hy_v4",
271
+ "moe_intermediate_size": 2048,
272
+ "mtp_loss_factor": 0.1,
273
+ "n_group": 1,
274
+ "n_routed_experts": 256,
275
+ "n_shared_experts": 1,
276
+ "norm_topk_prob": true,
277
+ "num_attention_heads": 64,
278
+ "num_experts_per_tok": 8,
279
+ "num_hidden_layers": 78,
280
+ "num_key_value_heads": 8,
281
+ "num_nextn_predict_layers": 1,
282
+ "pad_token_id": 120002,
283
+ "q_lora_rank": 2048,
284
+ "qk_head_dim": 256,
285
+ "qk_nope_head_dim": 192,
286
+ "qk_rope_head_dim": 64,
287
+ "rms_norm_eps": 1e-05,
288
+ "rope_parameters": {
289
+ "rope_theta": 10000000,
290
+ "rope_type": "default"
291
+ },
292
+ "routed_scaling_factor": 2.827,
293
+ "swiglu_limit": 10.0,
294
+ "tie_word_embeddings": false,
295
+ "topk_group": 1,
296
+ "transformers_version": "5.16.2",
297
+ "use_cache": true,
298
+ "use_dsa": true,
299
+ "use_mla": true,
300
+ "v_head_dim": 256,
301
+ "vocab_size": 120832,
302
+ "torch_dtype": "bfloat16"
303
+ }
finetune/README.md ADDED
@@ -0,0 +1,385 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ <p align="left">
2
+ <a href="README_CN.md">中文</a> | English
3
+ </p>
4
+
5
+ # Model Fine-tuning
6
+
7
+ Hy4 preview provides processes related to model fine-tuning. This section details how to process training data for model fine-tuning purposes.
8
+
9
+ ## Training Data Format and Processing
10
+
11
+ **Hy4 preview supports both "slow thinking" and "fast thinking" modes. You can control the mode via the `reasoning_effort` parameter (options: `high`, `no_think`).**
12
+
13
+ The training data should be formatted as a list of messages. By default, the system prompt for both training and inference is empty, but you may customize it as needed.
14
+
15
+ ```python
16
+ # Fast thinking pattern (no_think)
17
+ {"reasoning_effort": "no_think", "messages": [{"content": "You are a helpful assistant.\nThe current time is 2026-01-01 13:26:12 Thursday", "role": "system"}, {"content": "1+1=?", "role": "user"}, {"role": "assistant", "content": "1+1=2"}]}
18
+
19
+ # Slow thinking pattern (high)
20
+ {"reasoning_effort": "high", "messages": [{"content": "You are a helpful assistant.\nThe current time is 2026-01-01 13:26:12 Thursday", "role": "system"}, {"content": "1+1=?", "role": "user"}, {"role": "assistant", "content": "1+1=2", "reasoning_content": "The user is asking for the result of 1 + 1. In basic decimal arithmetic, 1 + 1 equals 2."}]}
21
+ ```
22
+
23
+ Example of using `apply_chat_template` to tokenize:
24
+
25
+ ```python
26
+ from transformers import AutoTokenizer
27
+ tokenizer = AutoTokenizer.from_pretrained("./models", use_fast=False, trust_remote_code=True)
28
+
29
+ messages = [
30
+ {"content": "You are a helpful assistant.", "role": "system"},
31
+ {"content": "1+1=?", "role": "user"},
32
+ {"role": "assistant", "content": "1+1=2"}
33
+ ]
34
+ ids = tokenizer.apply_chat_template(messages, tokenize=True, reasoning_effort="no_think")
35
+ ```
36
+
37
+ ## Fine-tuning Process
38
+
39
+ ### Hardware Requirements
40
+
41
+ Based on testing, the minimum resource configuration is as follows:
42
+
43
+ - **LoRA Fine-tuning**: At least 8 machines with 64 GPUs (at least 96GB GPU memory per GPU, at least 2TB CPU memory per machine).
44
+ - **Full Fine-tuning**: At least 16 machines with 128 GPUs (at least 96GB GPU memory per GPU, at least 2TB CPU memory per machine).
45
+
46
+ > Note: The above are minimum resource configurations; actual requirements increase with `max_seq_length`, batch size, etc.
47
+
48
+ ### Configure Passwordless SSH Login Between Machines (Multi-Machine Training)
49
+
50
+ > If you only use single-machine training, you can skip this section.
51
+
52
+ The following instructions use two machines as an example, with their IPs denoted as `${ip1}` and `${ip2}`. All steps should be performed inside the Docker container.
53
+
54
+ First, configure passwordless SSH for each container on every machine:
55
+
56
+ ```sh
57
+ ssh-keygen # Generate id_rsa and id_rsa.pub for passwordless login
58
+ ssh-keygen -t rsa -A # Generate /etc/ssh/ssh_host_rsa_key and ssh_host_ecdsa_key for SSH listening
59
+ /usr/sbin/sshd -p 36005 -o ListenAddress=0.0.0.0 # Start SSH listening
60
+ echo "Port 36005" > ~/.ssh/config # Set SSH connection port to 36005
61
+ passwd root # Set the root password to avoid monitoring platform alerts
62
+ ```
63
+
64
+ Note: `36005` is an example port. You may use any available port, but ensure it is **open** and **not occupied by other processes**.
65
+
66
+ Next, in each machine's container, execute:
67
+
68
+ ```sh
69
+ cat ~/.ssh/id_rsa.pub
70
+ ```
71
+
72
+ **Copy the output SSH public key and paste it into the `~/.ssh/authorized_keys` file, one key per line. This must be done on every machine.** In the end, the `~/.ssh/authorized_keys` file on each machine should be identical and contain the public keys of all machines.
73
+
74
+ Please note that for multi-node training, the code executed on each node must be identical. It is recommended to mount a shared network drive. If this is not possible, you must manually copy the dataset, scripts, and code to the same directory on each machine.
75
+
76
+ ### Launch Methods
77
+
78
+ This project provides three fine-tuning methods. You can choose based on your needs:
79
+
80
+ - **DeepSpeed Native Fine-tuning** (based on HuggingFace Transformers Trainer): Located in the `deepspeed_support` directory
81
+ - **LLaMA-Factory Fine-tuning**: Located in the `llama_factory_support` directory
82
+ - **ms-swift Fine-tuning**: Located in the `ms_swift_support` directory
83
+
84
+ #### DeepSpeed Native Fine-tuning
85
+
86
+ Reference: [HuggingFace Transformers Trainer](https://huggingface.co/docs/transformers/main/en/main_classes/trainer)
87
+
88
+ ##### Single-Machine Fine-tuning
89
+
90
+ In the `deepspeed_support` directory, execute:
91
+
92
+ ```sh
93
+ pip install -r requirements.txt
94
+ bash train.sh
95
+ ```
96
+
97
+ ##### Multi-Machine Fine-tuning
98
+
99
+ To launch fine-tuning across multiple machines, please first complete the configuration in [Configure Passwordless SSH Login Between Machines](#configure-passwordless-ssh-login-between-machines-multi-machine-training), and ensure all machines are within the same cluster.
100
+
101
+ Confirm that dependencies are installed (if not, run `pip install -r requirements.txt`), then add the following configuration at the beginning of `train.sh`:
102
+
103
+ ```shell
104
+ export HOST_GPU_NUM=8
105
+ # IP list, comma separated. e.g. "192.168.1.1,192.168.1.2" or single node "192.168.1.1"
106
+ IP_LIST=${IP_LIST:-"127.0.0.1"}
107
+ ```
108
+
109
+ Note: If the `IP_LIST` environment variable is not set, replace `IP_LIST` with the IP list! The format is:
110
+ ```
111
+ For a single IP:
112
+ IP_LIST=${ip_1}
113
+
114
+ For multiple IPs:
115
+ IP_LIST=${ip_1},${ip_2}
116
+
117
+ ```
118
+
119
+ Replace `${ip_1}` and `${ip_2}` with the actual IP addresses.
120
+
121
+ Then, on the machine with `${ip1}`, execute `bash train.sh` in the `deepspeed_support/` directory. On first launch, you may see the following output:
122
+
123
+ ```ssh
124
+ The authenticity of host '[ip]:36005 ([ip]:36005)' can't be established.
125
+ ECDSA key fingerprint is xxxxxx.
126
+ ECDSA key fingerprint is MD5:xxxxxx.
127
+ Are you sure you want to continue connecting (yes/no)?
128
+ ```
129
+
130
+ Type `yes` to continue.
131
+
132
+ ##### Key Parameters
133
+
134
+ The key parameters in the script are as follows:
135
+
136
+ - `--deepspeed`: Path to the DeepSpeed configuration file. Four default DeepSpeed configuration files are provided in the `deepspeed_support` folder: `ds_zero2_no_offload.json`, `ds_zero2_offload.json`, `ds_zero3_no_offload.json`, and `ds_zero3_offload.json`, with different ZeRO stages (ZeRO-2 / ZeRO-3) and offload strategies selectable based on available GPU memory and communication constraints.
137
+ - `--model_name_or_path`: Path to the Hy4 preview HF pre-trained model weights to load, otherwise loading will fail.
138
+ - `--tokenizer_name_or_path`: Path to the tokenizer folder, otherwise loading will fail.
139
+ - `--train_data_file`: Path to the training file, which should be a jsonl file.
140
+ - `--output_dir`: Output directory where logs, tensorboard files, and model weights will be stored.
141
+ - `--per_device_train_batch_size`: Batch size per GPU.
142
+ - `--gradient_accumulation_steps`: Number of gradient accumulation steps. The global batch size is `per_device_train_batch_size * gradient_accumulation_steps * dp_size`.
143
+ - `--max_steps`: Total number of training steps.
144
+ - `--save_steps`: Number of steps between saving checkpoints.
145
+ - `--use_lora`: Whether to use LoRA training. Also accepts `--lora_rank`, `--lora_alpha`, and `--lora_dropout` parameters. By default, LoRA is applied to "q_proj", "k_proj", "v_proj", and "o_proj". To change this, modify the code. Note: **When using LoRA training, only the LoRA weights are saved, not the base model weights.**
146
+ - `--make_moe_param_leaf_module`: When using ZeRO-3 with MoE training, treat the MoE module as a leaf module, i.e., its parameters are not partitioned by ZeRO-3. This option is expected to significantly increase memory usage.
147
+ - `--gradient_checkpointing`: Enable gradient checkpointing.
148
+ - `--learning_rate`: Maximum learning rate during training.
149
+ - `--min_lr`: Minimum learning rate during training.
150
+ - `--use_flash_attn`: Enable flash-attention for accelerated training.
151
+
152
+ **Notes:**
153
+
154
+ - To resume training from a previously saved checkpoint rather than loading pre-trained weights, specify `--resume_from_checkpoint` with the path to the checkpoint. Do not specify `--model_name_or_path`; this will load only the weights without the training state.
155
+ - When resuming from a checkpoint, there may be minor differences in loss due to the randomness of some non-deterministic algorithms. This is normal. See: [HuggingFace Transformers Trainer Randomness](https://huggingface.co/docs/transformers/main/en/main_classes/trainer#randomness)
156
+ - When `--model_name_or_path` is specified, all model-related parameters will be ignored.
157
+ - Samples within a batch are padded to the length of the longest sample in the batch, but the maximum length of each sample is `max_seq_length`. Any excess will be truncated.
158
+ - If you see a warning about **linear layer** bias weights not being loaded, you can ignore it; Hy4 preview's linear layers (q_proj / k_proj / v_proj / o_proj, etc.) do not use bias. Note: the MoE router's `e_score_correction_bias` is a buffer and is auto-loaded by the training script, so please do not ignore its loading failure.
159
+
160
+ ##### What if GPU Memory is Insufficient?
161
+
162
+ Reference: [DeepSpeed Configuration](https://www.deepspeed.ai/docs/config-json/)
163
+
164
+ You can try modifying the DeepSpeed configuration by removing the `auto` attribute from the following parameters and reducing their values:
165
+
166
+ - `stage3_param_persistence_threshold`
167
+ - `stage3_prefetch_bucket_size`
168
+ - `stage3_max_reuse_distance`
169
+
170
+ #### LLaMA-Factory Fine-tuning
171
+
172
+ If you are familiar with LLaMA-Factory, you may use it for fine-tuning. All scripts, code, and configuration files are archived in the `llama_factory_support` directory. Unless otherwise specified, all files mentioned below are located in this directory.
173
+
174
+ ##### Installation
175
+
176
+ You can install LLaMA-Factory by downloading the source code from https://github.com/hiyouga/LLaMA-Factory/tree/main and following the instructions on the website.
177
+
178
+ ##### Configuration Files
179
+
180
+ We provide sample LLaMA-Factory fine-tuning configuration files: `hy_v4_lora_sft.yaml` and `hy_v4_full_sft.yaml`, corresponding to LoRA fine-tuning and full fine-tuning respectively.
181
+
182
+ Key parameters in the configuration files are as follows:
183
+
184
+ **Model:**
185
+
186
+ - `model_name_or_path`: Path to the Hy4 preview HF format pre-trained model weights
187
+ - `trust_remote_code`: Whether to trust remote code; Hy4 preview requires this to be set to `true`
188
+
189
+ **Training Method:**
190
+
191
+ - `stage`: Training stage, currently `sft` (supervised fine-tuning)
192
+ - `finetuning_type`: Fine-tuning type, either `full` (full fine-tuning) or `lora` (LoRA fine-tuning)
193
+ - `deepspeed`: DeepSpeed configuration file path; `../deepspeed_support/ds_zero3_offload.json` is recommended for full fine-tuning
194
+ - `fsdp` + `fsdp_config`: FSDP distributed strategy; recommended for LoRA fine-tuning (configuration is built into `hy_v4_lora_sft.yaml`); mutually exclusive with DeepSpeed
195
+
196
+ > **Distributed Strategy Recommendations:**
197
+ > - **FSDP**: Recommended for LoRA fine-tuning, good compatibility and simple configuration
198
+ > - **DeepSpeed ZeRO-3 + Offload**: Recommended for full fine-tuning or memory-constrained scenarios
199
+
200
+ **LoRA Parameters (only effective during LoRA fine-tuning):**
201
+
202
+ - `lora_rank`: LoRA rank, default `64`
203
+ - `lora_alpha`: LoRA alpha coefficient, default `128`
204
+ - `lora_dropout`: LoRA dropout ratio, default `0.05`
205
+ - `lora_target`: Target modules for LoRA, default `q_a_proj,q_b_proj,kv_a_proj_with_mqa,kv_b_proj,o_proj`
206
+
207
+ **Dataset:**
208
+
209
+ - `dataset_dir`: Dataset directory path
210
+ - `dataset`: Dataset name, must be registered in `dataset_info.json` under `dataset_dir`
211
+ - `template`: Chat template; Hy4 preview uses `hy_v4`
212
+ - `cutoff_len`: Maximum sequence length; sequences exceeding this will be truncated. For LoRA fine-tuning, a smaller value is recommended to save memory
213
+ - `max_samples`: Maximum number of samples per dataset
214
+ - `overwrite_cache`: Whether to overwrite cached preprocessed datasets
215
+
216
+ **Output:**
217
+
218
+ - `output_dir`: Output directory where logs, TensorBoard files, and weights will be stored
219
+ - `logging_steps`: Number of steps between logging
220
+ - `save_steps`: Number of steps between saving checkpoints
221
+ - `plot_loss`: Whether to plot the training loss curve
222
+ - `overwrite_output_dir`: Whether to overwrite the existing output directory
223
+ - `save_only_model`: Whether to save only model weights (excluding optimizer states, etc.)
224
+ - `report_to`: Logging tool, options: `none`, `wandb`, `tensorboard`, `swanlab`, `mlflow`
225
+
226
+ **Training Hyperparameters:**
227
+
228
+ - `per_device_train_batch_size`: Batch size per GPU
229
+ - `gradient_accumulation_steps`: Gradient accumulation steps; `per_device_train_batch_size * gradient_accumulation_steps * dp_size` equals the global batch size
230
+ - `learning_rate`: Maximum learning rate; `1.0e-5` recommended for full fine-tuning, `2.0e-4` for LoRA fine-tuning
231
+ - `num_train_epochs`: Number of training epochs
232
+ - `lr_scheduler_type`: Learning rate scheduler type; `cosine_with_min_lr` is recommended
233
+ - `lr_scheduler_kwargs.min_lr_rate`: Ratio of minimum to maximum learning rate; e.g., `0.1` means the minimum learning rate is 10% of the maximum
234
+ - `warmup_steps`: Number of warmup steps
235
+ - `bf16`: Whether to use BFloat16 mixed precision training
236
+ - `gradient_checkpointing`: Whether to enable gradient checkpointing to save memory
237
+ - `ddp_timeout`: Distributed training timeout (milliseconds)
238
+ - `flash_attn`: Attention implementation; `auto` (automatic selection) or `sdpa` is recommended
239
+ - `resume_from_checkpoint`: Resume training from a specified checkpoint path; set to `null` to start from scratch
240
+
241
+ ##### Launch Fine-tuning
242
+
243
+ For multi-machine fine-tuning, please first complete the configuration in [Configure Passwordless SSH Login Between Machines](#configure-passwordless-ssh-login-between-machines-multi-machine-training) (single-machine fine-tuning can skip this step).
244
+
245
+ Modify the following configuration at the beginning of `train_lf.sh`:
246
+
247
+ ```shell
248
+ export HOST_GPU_NUM=8
249
+ # IP list, comma separated. e.g. "192.168.1.1,192.168.1.2" or single node "192.168.1.1"
250
+ export IP_LIST=${IP_LIST:-"127.0.0.1"}
251
+ ```
252
+
253
+ Note:
254
+ 1. If the `IP_LIST` environment variable is not set, replace `IP_LIST` with the IP list! The format is:
255
+ ```
256
+ For a single IP:
257
+ IP_LIST=${ip_1}
258
+
259
+ For multiple IPs:
260
+ IP_LIST=${ip_1},${ip_2}
261
+
262
+ ```
263
+ Replace `${ip_1}` and `${ip_2}` with the actual IP addresses.
264
+
265
+ 2. To specify a fine-tuning configuration file, set the `YAML_FILE` environment variable. The default is `hy_v4_full_sft.yaml`. For example, to use the LoRA fine-tuning configuration:
266
+ ```shell
267
+ export YAML_FILE=hy_v4_lora_sft.yaml
268
+ ```
269
+
270
+ Then, on each machine, run the launch script in the `llama_factory_support/` directory:
271
+
272
+ ```shell
273
+ bash train_lf.sh
274
+ ```
275
+
276
+ #### ms-swift Fine-tuning
277
+
278
+ If you are familiar with ms-swift, you may use it for fine-tuning. All scripts, code, and configuration files are archived in the `ms_swift_support` directory. Unless otherwise specified, all files mentioned below are located in this directory.
279
+
280
+ ##### Installation
281
+
282
+ You can install ms-swift via pip:
283
+
284
+ ```sh
285
+ pip install ms-swift
286
+ ```
287
+
288
+ Or install from source: https://github.com/modelscope/ms-swift
289
+
290
+ ##### Fine-tuning Scripts and Configuration Files
291
+
292
+ | Fine-tuning Method | Configuration File | Launch Script |
293
+ |----------------|-------------------|---------------|
294
+ | Full Fine-tuning | `hy_v4_full_sft.yaml` | `bash sft_train.sh` |
295
+ | LoRA Fine-tuning | `hy_v4_lora_sft.yaml` | `bash sft_train_lora.sh` |
296
+
297
+ ##### About the eos_token_id Patch
298
+
299
+ The `hy_v4_swift_patches.py` file in the directory fixes an issue with the eos token in ms-swift's default template. The default template uses the `<|hy_eos|>` string as `chat_sep` and `suffix`, which gets tokenized into multiple token IDs, causing `model.generate()` to fail to stop correctly during inference.
300
+
301
+ The patch re-registers the template using the `[['eos_token_id']]` syntax, allowing ms-swift to dynamically resolve `tokenizer.eos_token_id` at runtime and generate the correct single token.
302
+
303
+ The launch script automatically loads this patch via `--custom_register_path hy_v4_swift_patches.py`, requiring no additional action.
304
+
305
+ ##### Key Parameters
306
+
307
+ Key parameters in the configuration files are as follows:
308
+
309
+ **Model:**
310
+
311
+ - `model`: Model path, can be a HuggingFace Hub ID or a local path
312
+ - `model_type`: Model type, set to `hy_v4`
313
+ - `template`: Chat template, set to `hy_v4`
314
+ - `torch_dtype`: Data type, `bfloat16` is recommended
315
+ - `attn_impl`: Attention implementation, `sdpa` is recommended
316
+
317
+ **Training Method:**
318
+
319
+ - `train_type`: Fine-tuning type; set to `full` for full fine-tuning, `lora` for LoRA fine-tuning
320
+ - `lora_rank`: LoRA rank, default `64`
321
+ - `lora_alpha`: LoRA alpha coefficient, default `128`
322
+ - `lora_dropout`: LoRA dropout ratio, default `0.05`
323
+
324
+ **Dataset:**
325
+
326
+ - `dataset`: Dataset path, supports local jsonl files (sharegpt format)
327
+ - `max_length`: Maximum sequence length; sequences exceeding this will be truncated
328
+ - `lazy_tokenize`: Whether to use lazy tokenization, `true` is recommended
329
+
330
+ **Output:**
331
+
332
+ - `output_dir`: Output directory
333
+ - `save_steps`: Number of steps between saving checkpoints
334
+ - `save_total_limit`: Maximum number of checkpoints to keep
335
+ - `logging_steps`: Number of steps between logging
336
+ - `report_to`: Logging tool, options: `none`, `wandb`, `tensorboard`, `swanlab`, `mlflow`
337
+
338
+ **Training Hyperparameters:**
339
+
340
+ - `per_device_train_batch_size`: Batch size per GPU
341
+ - `gradient_accumulation_steps`: Gradient accumulation steps
342
+ - `learning_rate`: Maximum learning rate; `1.0e-5` recommended for full fine-tuning, `2.0e-4` for LoRA fine-tuning
343
+ - `num_train_epochs`: Number of training epochs
344
+ - `lr_scheduler_type`: Learning rate scheduler type, `cosine` is recommended
345
+ - `warmup_steps`: Number of warmup steps
346
+ - `bf16`: Whether to use BFloat16 mixed precision training
347
+
348
+ **Distributed Strategy / Optimization:**
349
+
350
+ - `deepspeed`: DeepSpeed strategy, options: `zero0`, `zero2`, `zero2_offload`, `zero3`, `zero3_offload`; `zero3_offload` recommended for full fine-tuning
351
+ - `fsdp` + `fsdp_config`: FSDP distributed strategy; recommended for LoRA fine-tuning; mutually exclusive with DeepSpeed
352
+ - `gradient_checkpointing`: Whether to enable gradient checkpointing
353
+ - `max_grad_norm`: Gradient clipping threshold
354
+
355
+ > **Distributed Strategy Recommendations:**
356
+ > - **FSDP**: Recommended for LoRA fine-tuning, good compatibility and simple configuration
357
+ > - **DeepSpeed ZeRO-3 + Offload**: Recommended for full fine-tuning or memory-constrained scenarios
358
+
359
+ **Other:**
360
+
361
+ - `ddp_timeout`: Distributed training timeout (milliseconds)
362
+ - `seed`: Random seed
363
+ - `resume_from_checkpoint`: Resume training from a specified checkpoint path
364
+
365
+ ##### Launch Fine-tuning
366
+
367
+ For multi-machine fine-tuning, please first complete the configuration in [Configure Passwordless SSH Login Between Machines](#configure-passwordless-ssh-login-between-machines-multi-machine-training) (single-machine fine-tuning can skip this step).
368
+
369
+ Modify the following configuration in the `sft_train.sh` script:
370
+
371
+ ```shell
372
+ export HOST_GPU_NUM=8
373
+ # IP list, comma separated. e.g. "10.0.0.1,10.0.0.2" or single node "127.0.0.1"
374
+ export IP_LIST=${IP_LIST:-"127.0.0.1"}
375
+ ```
376
+
377
+ Then, on each machine, execute the launch script in the `ms_swift_support/` directory:
378
+
379
+ ```sh
380
+ # Single-machine training
381
+ bash sft_train.sh
382
+
383
+ # Multi-machine training (execute on each machine)
384
+ IP_LIST="10.0.0.1,10.0.0.2" bash sft_train.sh
385
+ ```
finetune/README_CN.md ADDED
@@ -0,0 +1,384 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ <p align="left">
2
+ <a href="README.md">English</a> | 中文
3
+ </p>
4
+
5
+ # 模型微调
6
+
7
+ Hy4 preview 提供了模型微调相关流程,您可以在此章节对训练数据格式进行处理以供模型微调使用。
8
+
9
+ ## 训练数据格式及处理
10
+
11
+ **Hy4 preview 同时支持慢思考与快思考两种模式,可通过 `reasoning_effort` 参数控制(可选值:`high`、`no_think`)。**
12
+
13
+ 训练数据按照以下形式处理为 messages 格式,训练和推理的默认 system prompt 为空,可以根据自己的需求进行设定。
14
+
15
+ ```python
16
+ # Fast thinking pattern (no_think)
17
+ {"reasoning_effort": "no_think", "messages": [{"content": "你是一个有用的人工智能助手。\n现在的时间是2026-01-01 13:26:12 周四", "role": "system"}, {"content": "1+1=?", "role": "user"}, {"role": "assistant", "content": "1+1=2"}]}
18
+
19
+ # Slow thinking pattern (high)
20
+ {"reasoning_effort": "high", "messages": [{"content": "你是一个有用的人工智能助手。\n现在的时间是2026-01-01 13:26:12 周四", "role": "system"}, {"content": "1+1=?", "role": "user"}, {"role": "assistant", "content": "1+1=2", "reasoning_content": "用户问的是1+1等于多少。在基本的十进制算术中,1+1等于2。"}]}
21
+ ```
22
+
23
+ 使用 `apply_chat_template` 进行 tokenize 的示例:
24
+
25
+ ```python
26
+ from transformers import AutoTokenizer
27
+ tokenizer = AutoTokenizer.from_pretrained("./models", use_fast=False, trust_remote_code=True)
28
+
29
+ messages = [
30
+ {"content": "你是一个有用的人工智能助手。", "role": "system"},
31
+ {"content": "1+1=?", "role": "user"},
32
+ {"role": "assistant", "content": "1+1=2"}
33
+ ]
34
+ ids = tokenizer.apply_chat_template(messages, tokenize=True, reasoning_effort="no_think")
35
+ ```
36
+
37
+ ## 微调流程
38
+
39
+ ### 硬件需求
40
+
41
+ 经过测试,最小的资源配置如下:
42
+
43
+ - **LoRA 微调**:最少需要 8 机 64 卡,每卡显存至少 96GB,每机 CPU 内存至少 2TB。
44
+ - **全量微调**:最少需要 16 机 128 卡,每卡显存至少 96GB,每机 CPU 内存至少 2TB。
45
+
46
+ > 注:以上为最小资源配置,实际所需资源会随 max_seq_length、batch size 等因素相应增加。
47
+
48
+ ### 配置机器间免密 ssh 登录(多机训练)
49
+
50
+ > 如果只使用单机训练,可跳过本节。
51
+
52
+ 以下操作以两个机器为例,两台机器的 ip 分别以`${ip1}`和`${ip2}`标识,以下操作均在 docker container 内执行。
53
+
54
+ 首先,配置多机container免密,在每台机器上执行。
55
+
56
+ ```sh
57
+ ssh-keygen # 生成id_rsa和id_rsa.pub,用于免密登录
58
+ ssh-keygen -t rsa -A # 生成/etc/ssh/ssh_host_rsa_key和ssh_host_ecdsa_key, 用于后面启动ssh listen
59
+ /usr/sbin/sshd -p 36005 -o ListenAddress=0.0.0.0 # 启动 SSH 监听
60
+ echo "Port 36005" > ~/.ssh/config # ssh 连接端口修改为 36005
61
+ passwd root # 需要配置root密码,否则监测平台会报警
62
+ ```
63
+
64
+ 注意:这里的`36005`是一个示例端口,可以选用任意端口,但需要保证使用的端口**开放**且**不被其他的进程占用**。
65
+
66
+ 接下来,在每台机器的 container 内,执行:
67
+
68
+ ```sh
69
+ cat ~/.ssh/id_rsa.pub
70
+ ```
71
+
72
+ **将输出的 ssh 公钥复制并粘贴到`~/.ssh/authorized_keys`文件中,每行一个公钥,每台机器上都要做这个操作**。最终每台机器上的`~/.ssh/authorized_keys`文件内容应当是一致的,并且包含了所有机器的公钥。
73
+
74
+ 需要注意,多节点训练时,每个节点上执行的代码都得一致,建议挂载一个共享的网络盘,如果无法挂载共享网盘,则需要手动将数据集、脚本、代码复制在多台机器的相同目录下。
75
+
76
+ ### 启动方式
77
+
78
+ 本项目提供三种微调方式,您可以根据需求选择:
79
+
80
+ - **DeepSpeed 原生微调**(基于 HuggingFace Transformers Trainer):位于 `deepspeed_support` 目录下
81
+ - **LLaMA-Factory 微调**:位于 `llama_factory_support` 目录下
82
+ - **ms-swift 微调**:位于 `ms_swift_support` 目录下
83
+
84
+ #### DeepSpeed 原生微调
85
+
86
+ 参考:[HuggingFace Transformers Trainer](https://huggingface.co/docs/transformers/main/en/main_classes/trainer)
87
+
88
+ ##### 单机启动微调
89
+
90
+ 在 `deepspeed_support` 目录下,执行:
91
+
92
+ ```sh
93
+ pip install -r requirements.txt
94
+ bash train.sh
95
+ ```
96
+
97
+ ##### 多机启动微调
98
+
99
+ 如果要用多台机器启动微调,请先完成 [配置机器间免密 ssh 登录](#配置机器间免密-ssh-登录多机训练) 中的配置,并保证多台机器在一个集群内。
100
+
101
+ 确认依赖已经安装完成(如未安装,请执行`pip install -r requirements.txt`安装),然后在`train.sh`中的开头增加以下配置:
102
+
103
+ ```shell
104
+ export HOST_GPU_NUM=8
105
+ # IP list, comma separated. e.g. "192.168.1.1,192.168.1.2" or single node "192.168.1.1"
106
+ IP_LIST=${IP_LIST:-"127.0.0.1"}
107
+ ```
108
+
109
+ 注意:如果`IP_LIST`环境变量未设置,则将`IP_LIST`替换为IP列表!格式为:
110
+ ```
111
+ 如果只有一个IP:
112
+ IP_LIST=${ip_1}
113
+
114
+ 如果有多个IP:
115
+ IP_LIST=${ip_1},${ip_2}
116
+
117
+ ```
118
+
119
+ 请将`${ip_1}`和`${ip_2}`替换为真实的IP地址。
120
+
121
+ ��后,在`${ip1}`的机器上,在`deepspeed_support/`目录下,执行`bash train.sh`即可,注意第一次启动时可能会看见以下的输出:
122
+
123
+ ```ssh
124
+ The authenticity of host '[ip]:36005 ([ip]:36005)' can't be established.
125
+ ECDSA key fingerprint is xxxxxx.
126
+ ECDSA key fingerprint is MD5:xxxxxx.
127
+ Are you sure you want to continue connecting (yes/no)?
128
+ ```
129
+
130
+ 此时输入`yes`即可继续。
131
+
132
+ ##### 关键参数
133
+
134
+ 脚本中的关键参数如下:
135
+
136
+ - `--deepspeed`: 此参数应当指向一个 deepspeed 的配置文件,`deepspeed_support`文件夹下提供了四种 DeepSpeed 的默认配置文件:`ds_zero2_no_offload.json`, `ds_zero2_offload.json`, `ds_zero3_no_offload.json`, `ds_zero3_offload.json`,可根据显存与通信情况选择不同的 ZeRO 阶段(ZeRO-2 / ZeRO-3)与 offload 策略
137
+ - `--model_name_or_path`: 要加载的 Hy4 preview 的 HF 预训练模型权重,否则无法加载
138
+ - `--tokenizer_name_or_path`: tokenizer 文件夹路径, 否则无法加载
139
+ - `--train_data_file`: 训练文件路径,应该为一个 jsonl 文件
140
+ - `--output_dir`: 输出文件夹,log、tensorboard 和权重都会存储在这个路径下
141
+ - `--per_device_train_batch_size`: 每张卡上的 batch size
142
+ - `--gradient_accumulation_steps`: 梯度累计次数,`per_device_train_batch_size * gradient_accumulation_steps * dp_size`为 global_batch_size
143
+ - `--max_steps`: 训练的总步数
144
+ - `--save_steps`: 每多少个 step 存储一个 checkpoint
145
+ - `--use_lora`: 是否用 lora 训练,同时接收`--lora_rank`,`--lora_alpha`和`--lora_dropout`参数。lora 默认应用于 MLA(Multi-head Latent Attention)投影层:"q_a_proj", "q_b_proj", "kv_a_proj_with_mqa", "kv_b_proj", "o_proj",如果需要改变的话在代码中修改即可。注意:**使用 lora 训练时,只会保存 lora 的权重,而不会保存 base 模型的权重**
146
+ - `--make_moe_param_leaf_module`:当用 zero3 以及 MoE 训练时,将 MoE 模块视作一个 leaf module,即它的参数不进行 zero3 切分,这个选项预计会显著增加显存占用
147
+ - `--gradient_checkpointing`:开启梯度检查点- `--learning_rate`: 训练时的最大学习率
148
+ - `--min_lr`: 训练时的最小学习率
149
+ - `--use_flash_attn`: 开启 flash-attention 进行训练加速
150
+
151
+ **注意:**
152
+
153
+ - 如果想从一个中途保存的 ckpt 继续训练,而不是加载一个预训练的权重,直接指定`--resume_from_checkpoint`为之前训练保存的 ckpt 路径,不要指定`--model_name_or_path`,这样只会加载权重,而不会加载训练状态
154
+ - 从 ckpt 继续训练时,loss 可能会有微小的偏差,这是由一些非确定性算法带来的随机性,是正常现象。参考:[HuggingFace Transformers Trainer Randomness](https://huggingface.co/docs/transformers/main/en/main_classes/trainer#randomness)
155
+ - 当 `--model_name_or_path` 有效时,所有模型相关的参数都会被忽略
156
+ - 一个 batch 内的样本会通过 padding 对齐 batch 内最长的样本,而每条样本的长度最长为 max_seq_length,超出的部分会被裁剪
157
+ - 如果报出**线性层** bias 权重没有 load 的 warning,忽略即可,Hy4 preview 的线性层(q_a_proj / q_b_proj / kv_a_proj_with_mqa / kv_b_proj / o_proj 等)不使用 bias。注意:MoE 路由的 `e_score_correction_bias` 属于 buffer,已由训练脚本自动加载,如果加载失败请不要忽略。
158
+
159
+ ##### 显存不足怎么办?
160
+
161
+ 参考:[DeepSpeed Configuration](https://www.deepspeed.ai/docs/config-json/)
162
+
163
+ 可以尝试修改 ds config,去掉这几个参数的 auto 属性,改小试试看:
164
+
165
+ - `stage3_param_persistence_threshold`
166
+ - `stage3_prefetch_bucket_size`
167
+ - `stage3_max_reuse_distance`
168
+
169
+ #### LLaMA-Factory 微调
170
+
171
+ 如果对 LLaMA-Factory 较为熟悉,可使用 LLaMA-Factory 进行微调。脚本、代码以及配置文件都归档在 `llama_factory_support` 目录下。如果没有特别说明,接下来我们提到的文件都是该目录下的文件。
172
+
173
+ ##### 安装
174
+
175
+ 可以通过下载源码 https://github.com/hiyouga/LLaMA-Factory/tree/main ,根据网站的指引进行安装。
176
+
177
+ ##### 配置文件
178
+
179
+ 我们提供了 llama-factory 的微调示例配置文件 `hy_v4_lora_sft.yaml`和`hy_v4_full_sft.yaml`文件,分别对应 LoRA 微调和全量微调。
180
+
181
+ 脚本中的关键参数如下:
182
+
183
+ **模型相关:**
184
+
185
+ - `model_name_or_path`: Hy4 preview HF 格式预训练模型权重路径
186
+ - `trust_remote_code`: 是否信任远程代码, Hy4 preview 需要设置为 `true`
187
+
188
+ **训练方法:**
189
+
190
+ - `stage`: 训练阶段, 当前为 `sft`(监督微调)
191
+ - `finetuning_type`: 微调类型, 可选 `full`(全量微调) 或 `lora`(LoRA 微调)
192
+ - `deepspeed`: DeepSpeed 配置文件路径, 全量微调推荐 `../deepspeed_support/ds_zero3_offload.json`
193
+ - `fsdp` + `fsdp_config`: FSDP 分布式策略, LoRA 微调推荐使用 FSDP(配置已内置在 `hy_v4_lora_sft.yaml` 中); 与 DeepSpeed 二选一
194
+
195
+ > **分布式策略选择建议:**
196
+ > - **FSDP**:推荐用于 LoRA 微调,兼容性好,配置简单
197
+ > - **DeepSpeed ZeRO-3 + Offload**:推荐用于全量微调或显存紧张的场景
198
+
199
+ **LoRA 参数(仅 LoRA 微调时生效):**
200
+
201
+ - `lora_rank`: LoRA 秩, 默认 `64`
202
+ - `lora_alpha`: LoRA alpha 系数, 默认 `128`
203
+ - `lora_dropout`: LoRA dropout 比率, 默认 `0.05`
204
+ - `lora_target`: LoRA 应用的目标模块, 默认为 `q_a_proj,q_b_proj,kv_a_proj_with_mqa,kv_b_proj,o_proj`
205
+
206
+ **数据集:**
207
+
208
+ - `dataset_dir`: 数据集目录路径
209
+ - `dataset`: 数据集名称, 需要在 `dataset_dir` 下的 `dataset_info.json` 中注册
210
+ - `template`: 对话模板, Hy4 preview 使用 `hy_v4`
211
+ - `cutoff_len`: 最大序列长度, 超出部分会被截断; LoRA 微调建议适当减小以节省显存
212
+ - `max_samples`: 每个数据集最多使用的样本数
213
+ - `overwrite_cache`: 是否覆盖已缓存的预处理数据集
214
+
215
+ **输出:**
216
+
217
+ - `output_dir`: 输出目录, 日志、TensorBoard 和权重都会存储在此路径下
218
+ - `logging_steps`: 每多少步记录一次日志
219
+ - `save_steps`: 每多少步保存一次 checkpoint
220
+ - `plot_loss`: 是否绘制训练 loss 曲线
221
+ - `overwrite_output_dir`: 是否覆盖已有的输出目录
222
+ - `save_only_model`: 是否只保存模型权重(不保存优化器状态等)
223
+ - `report_to`: 日志上报工具, 可选 `none`, `wandb`, `tensorboard`, `swanlab`, `mlflow`
224
+
225
+ **训练超参数:**
226
+
227
+ - `per_device_train_batch_size`: 每张卡上的 batch size
228
+ - `gradient_accumulation_steps`: 梯度累积步数, `per_device_train_batch_size * gradient_accumulation_steps * dp_size` 为 global batch size
229
+ - `learning_rate`: 最大学习率, 全量微调推荐 `1.0e-5`, LoRA 微调推荐 `2.0e-4`
230
+ - `num_train_epochs`: 训练轮数
231
+ - `lr_scheduler_type`: 学习率调度器类型, 推荐使用 `cosine_with_min_lr`
232
+ - `lr_scheduler_kwargs.min_lr_rate`: 最小学习率与最大学习率的比值, 例如 `0.1` 表示最小学习率为最大学习率的 10%
233
+ - `warmup_steps`: 预热步数
234
+ - `bf16`: 是否使用 BFloat16 混合精度训练
235
+ - `gradient_checkpointing`: 是否开启梯度重计算以节省显存
236
+ - `ddp_timeout`: 分布式训练超时时间(毫秒)
237
+ - `flash_attn`: 注意力实现方式, 推荐 `auto`(自动选择) 或 `sdpa`
238
+ - `resume_from_checkpoint`: 从指定 checkpoint 路径恢复训练, 设为 `null` 表示从头开始训练
239
+
240
+ ##### 启动微调
241
+
242
+ 如需多机训练,请先完成 [配置机器间免密 ssh 登录](#配置机器间免密-ssh-登录多机训练) 中的配置(单机训练可跳过此步骤)。
243
+
244
+ 修改`train_lf.sh`中开头的以下配置:
245
+
246
+ ```shell
247
+ export HOST_GPU_NUM=8
248
+ # IP list, comma separated. e.g. "192.168.1.1,192.168.1.2" or single node "192.168.1.1"
249
+ export IP_LIST=${IP_LIST:-"127.0.0.1"}
250
+ ```
251
+
252
+ 注意:
253
+ 1. 如果`IP_LIST`环境变量未设置,则将`IP_LIST`替换为IP列表!格式为:
254
+ ```
255
+ 如果只有一个IP:
256
+ IP_LIST=${ip_1}
257
+
258
+ 如果有多个IP:
259
+ IP_LIST=${ip_1},${ip_2}
260
+
261
+ ```
262
+ 请将`${ip_1}`和`${ip_2}`替换为真实的IP地址。
263
+
264
+ 2. 如需指定微调配置文件,可通过`YAML_FILE`环境变量设置,默认为`hy_v4_full_sft.yaml`。例如使用 LoRA 微调配置:
265
+ ```shell
266
+ export YAML_FILE=hy_v4_lora_sft.yaml
267
+ ```
268
+
269
+ 然后,在每一台机器上,在`llama_factory_support/`目录下执行启动脚本:
270
+
271
+ ```shell
272
+ bash train_lf.sh
273
+ ```
274
+
275
+ #### ms-swift 微调
276
+
277
+ 如果对 ms-swift 较为熟悉,可使用 ms-swift 进行微调。脚本、代码以及配置文件都归档在 `ms_swift_support` 目录下。如果没有特别说明,接下来我们提到的文件都是该目录下的文件。
278
+
279
+ ##### 安装
280
+
281
+ 可以通过 pip 安装 ms-swift:
282
+
283
+ ```sh
284
+ pip install ms-swift
285
+ ```
286
+
287
+ 或从源码安装:https://github.com/modelscope/ms-swift
288
+
289
+ ##### 微调脚本与配置文件
290
+
291
+ | 微调方式 | 配置文件 | 启动脚本 |
292
+ |---------|---------|---------|
293
+ | 全量微调 | `hy_v4_full_sft.yaml` | `bash sft_train.sh` |
294
+ | LoRA 微调 | `hy_v4_lora_sft.yaml` | `bash sft_train_lora.sh` |
295
+
296
+ ##### 关于 eos_token_id Patch
297
+
298
+ 目录下的 `hy_v4_swift_patches.py` 文件用于修复 ms-swift 默认模板中 eos token 的问题。默认模板将 `<|hy_eos|>` 字符串作为 `chat_sep` 和 `suffix`,该字符串会被 tokenize 为多个 token ID,导致推理时 `model.generate()` 无法正确停止。
299
+
300
+ Patch 通过 `[['eos_token_id']]` 语法重新注册模板,使 ms-swift 在运行时动态解析 `tokenizer.eos_token_id`,生成正确的单个 token。
301
+
302
+ 启动脚本已通过 `--custom_register_path hy_v4_swift_patches.py` 自动加载此 patch,无需额外操作。
303
+
304
+ ##### 关键参数
305
+
306
+ 配置文件中的关键参数如下:
307
+
308
+ **模型相关:**
309
+
310
+ - `model`: 模型路径,可以是 HuggingFace Hub ID 或本地路径
311
+ - `model_type`: 模型类型,设为 `hy_v4`
312
+ - `template`: 对话模板,设为 `hy_v4`
313
+ - `torch_dtype`: 数据类型,推荐 `bfloat16`
314
+ - `attn_impl`: 注意力实现,推荐 `sdpa`
315
+
316
+ **训练方法:**
317
+
318
+ - `train_type`: 微调类型,全量微调设为 `full`,LoRA 微调设为 `lora`
319
+ - `lora_rank`: LoRA 秩,默认 `64`
320
+ - `lora_alpha`: LoRA alpha 系数,默认 `128`
321
+ - `lora_dropout`: LoRA dropout 比率,默认 `0.05`
322
+
323
+ **数据集:**
324
+
325
+ - `dataset`: 数据集路径,支持本地 jsonl 文件(sharegpt 格式)
326
+ - `max_length`: 最大序列长度,超出部分会被截断
327
+ - `lazy_tokenize`: 是否延迟 tokenize,推荐 `true`
328
+
329
+ **输出:**
330
+
331
+ - `output_dir`: 输出目录
332
+ - `save_steps`: 每多少步保存一次 checkpoint
333
+ - `save_total_limit`: 最多保留的 checkpoint 数量
334
+ - `logging_steps`: 每多少步记录一次日志
335
+ - `report_to`: 日志上报工具,可选 `none`, `wandb`, `tensorboard`, `swanlab`, `mlflow`
336
+
337
+ **训练超参数:**
338
+
339
+ - `per_device_train_batch_size`: 每张卡上的 batch size
340
+ - `gradient_accumulation_steps`: 梯度累积步数
341
+ - `learning_rate`: 最大学习率,全量微调推荐 `1.0e-5`,LoRA 微调推荐 `2.0e-4`
342
+ - `num_train_epochs`: 训练轮数
343
+ - `lr_scheduler_type`: 学习率调度器类型,推荐 `cosine`
344
+ - `warmup_steps`: 预热步数
345
+ - `bf16`: 是否使用 BFloat16 混合精度训练
346
+
347
+ **分布式策略 / 优化:**
348
+
349
+ - `deepspeed`: DeepSpeed 策略,可选 `zero0`, `zero2`, `zero2_offload`, `zero3`, `zero3_offload`;全量微调推荐 `zero3_offload`
350
+ - `fsdp` + `fsdp_config`: FSDP 分布式策略,LoRA 微调推荐使用 FSDP;与 DeepSpeed 二选一
351
+ - `gradient_checkpointing`: 是否开启梯度重计算
352
+ - `max_grad_norm`: 梯度裁剪阈值
353
+
354
+ > **分布式策略选择建议:**
355
+ > - **FSDP**:推荐用于 LoRA 微调,兼容性好,配置简单
356
+ > - **DeepSpeed ZeRO-3 + Offload**:推荐用于全量微调或显存紧张的场景
357
+
358
+ **其他:**
359
+
360
+ - `ddp_timeout`: 分布式训练超时时间(毫秒)
361
+ - `seed`: 随机种子
362
+ - `resume_from_checkpoint`: 从指定 checkpoint 路径恢复训练
363
+
364
+ ##### 启动微调
365
+
366
+ 如需多机训练,请先完成 [配置机器间免密 ssh 登录](#配置机器间免密-ssh-登录多机训练) 中的配置(单机训练可跳过此步骤)。
367
+
368
+ 修改 `sft_train.sh` 脚本中的以下配置:
369
+
370
+ ```shell
371
+ export HOST_GPU_NUM=8
372
+ # IP list, comma separated. e.g. "10.0.0.1,10.0.0.2" or single node "127.0.0.1"
373
+ export IP_LIST=${IP_LIST:-"127.0.0.1"}
374
+ ```
375
+
376
+ 然后,在每一台机器上,在 `ms_swift_support/` 目录下执行启动脚本:
377
+
378
+ ```sh
379
+ # 单机训练
380
+ bash sft_train.sh
381
+
382
+ # 多机训练(在每台机器上执行)
383
+ IP_LIST="10.0.0.1,10.0.0.2" bash sft_train.sh
384
+ ```
finetune/data/example_data.jsonl ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ {"messages": [{"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "能否帮我撰写一个关于环保的议论文开头部分?\n主题:环保\n立场:重要性和我们的责任"}, {"role": "assistant", "content": "在我们的生活中,环保已经成为一个不可忽视的议题。随着科技的发展和人口的增长,我们的环境正在遭受前所未有的压力。森林被砍伐,河流被污染,动植物种类正在消失,全球气候正在发生变化。这些都是我们无法忽视的事实。然而,环保并不仅仅是一个环境问题,它更是一个关乎我们生存和发展的问题。环保的重要性已经超越了任何一个单一的领域,它关乎我们的健康、经济、社会和文化。因此,我们每一个人都有责任和义务去保护我们的环境,这不仅是为了我们自己,也是为了我们的后代和整个地球。"}]}
2
+ {"messages": [{"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "请描述小红书APP的用户画像,并用表格的形式展示结果。"}, {"role": "assistant", "content": "| 用户特征 | 描述 |\n| -------- | ------------------------------ |\n| 年龄段 | 18-35岁 |\n| 性别 | 女性用户较多(约占70%) |\n| 地域 | 一线城市和二线城市用户占比较高 |\n| 职业 | 白领、创业者、自由职业者等 |\n| 爱好 | 美容、时尚、旅行、美食、健康、生活方式等 |\n| 消费能力 | 中高消费能力用户 |\n| 使用目的 | 阅读内容,分享购物心得,结识相似兴趣的朋友 |\n\n注:以上用户画像仅供参考,实际情况可能会有所不同。"}]}
3
+ {"messages": [{"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "请根据以下两个句子的描述,判断它们是否对同一情况进行了描述。句子1:“外面下着大雨,人们纷纷打起了伞。” 句子2:“由于雨量较大,人们都撑起了雨伞。”"}, {"role": "assistant", "content": "是的,这两个句子描述的是同一情况。"}]}
4
+ {"messages": [{"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "你认为游戏开发者应该如何处理游戏中的社交网络集成?"}, {"role": "assistant", "content": "游戏开发者在处理游戏中的社交网络集成时可以考虑以下几点:\n\n1. 以用户为中心:确保集成社交网络功能是为了提高用户体验,让玩家更容易地与朋友互动、分享成就和游戏进度。\n\n2. 保护隐私:在集成社交网络时,务必保护用户的隐私。在收集和使用用户数据时,需要遵循相关法规和用户协议,同时提供明确的隐私政策。\n\n3. 适度推送:在推送游戏动态、邀请和通知时要适度,避免过度打扰用户或侵犯用户隐私。可提供自定义设置,让玩家选择性地接收通知。\n\n4. 社交功能设计:在游戏中加入社交功能,如聊天、组队、比赛等,以增进玩家之间的互动。同时,确保这些功能易用且不影响游戏体验。\n\n5. 社区建设:鼓励玩家在社交平台上建立游戏相关的社区,分享游戏心得、攻略和经验。这将有助于增加游戏的知名度和吸引力。\n\n6. 奖励机制:设定奖励机制,鼓励玩家在社交网络上分享游戏内容、邀请好友等。但要确保奖励机制不会破坏游戏平衡。\n\n7. 持续更新和优化:根据玩家的反馈,不断优化社交网络集成功能,确保其与游戏的其他方面保持一致。\n\n总之,游戏开发者应该在保护用户隐私的前提下,提供便捷的社交网络集成功能,以提高游戏的互动性和吸引力。同时,要关注玩家的反馈,不断优化功能。"}]}
5
+ {"messages": [{"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "将这段过去进行时的文字改为现在完成进行时:“她们在公园里跳舞。”"}, {"role": "assistant", "content": "她们已经在公园里跳舞了。"}]}
6
+ {"messages": [{"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "在godaddy上买了一台服务器怎么更改root账户"}, {"role": "assistant", "content": "在GoDaddy上购买的服务器可能是Linux系统的VPS或者是专用服务器。要更改root账户,你可以通过SSH连接到服务器并修改密码。以下是具体步骤:\n\n1. 打开终端(Terminal)或使用SSH客户端(如PuTTY)。\n2. 使用SSH连接到服务器。在终端中输入以下命令:\n\n ```\n ssh root@your_server_ip_address\n ```\n\n 将`your_server_ip_address`替换为你的服务器IP地址。如果需要,你可能还需要输入服务器的SSH端口号。\n \n3. 当系统提示输入密码时,输入你的当前root密码。如果连接成功,你将看到服务器的命令提示符。\n\n4. 要更改root密码,请输入以下命令:\n\n ```\n passwd\n ```\n\n5. 系统将提示你输入新的UNIX密码。输入你想要设置的新密码并按回车键。\n\n6. 系统将要求你重新输入新密码以进行确认。再次输入新密码并按回车键。\n\n7. 如果密码更改成功,系统将显示“密码已成功更新”的消息。\n\n8. 使用`exit`命令断开与服务器的连接。\n\n现在你已经成功更改了root账户的密码。下次使用SSH连接到服务器时,请使用新的密码进行身份验证。"}]}
7
+ {"messages": [{"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "PHP算法逻辑:如何计算年龄?"}, {"role": "assistant", "content": "要计算年龄,您可以使用以下PHP算法。此算法考虑了生日、当前日期以及闰年的计算。\n\n```php\nfunction calculateAge($birthdate) {\n // 格式为YYYY-MM-DD的生日\n $birthdate = strtotime($birthdate);\n \n // 获取当前日期\n $now = time();\n\n // 计算年龄\n $age = date('Y', $now) - date('Y', $birthdate);\n \n // 如果还没有过生日,减去一年\n if (date('m-d', $now) < date('m-d', $birthdate)) {\n $age--;\n }\n\n return $age;\n}\n\n// 使用\n$birthdate = \"2000-01-01\";\necho \"年龄是: \" . calculateAge($birthdate) . \" 岁\";\n```\n\n这个函数将生日作为输入,输出用户的年龄。首先,它将生日转换为时间戳,然后获取当前的时间戳。接下来,它计算年份之间的差异。如果当前日期仍未超过这一年的生日,这个函数将年龄减少1。\n\n在使用函数时,只需传入出生日期即可,函数会返回相应的年龄。"}]}
8
+ {"messages": [{"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "给我解释一下什么是胰岛素抵抗? 请详细描述胰岛素抵抗的症状、原因与治疗方法。 胰岛素抵抗会影响一个人的日常生活吗? 有哪些生活习惯改变可以帮助改善胰岛素抵抗?"}, {"role": "assistant", "content": "胰岛素抵抗是指身体对胰岛素的敏感性降低,胰岛素不能有效地发挥作用,导致血糖无法正常进入身体细胞。胰岛素是一种激素,由胰腺产生,其主要作用是帮助身体细胞吸收血糖并将其转化为能量。当胰岛素抵抗发生时,胰腺会分泌更多的胰岛素以维持正常的血糖水平,但长期的胰岛素抵抗可能导致2型糖尿病和其他健康问题的发生。\n\n胰岛素抵抗的症状可能包括:\n\n1. 高血糖:由于胰岛素不能正常工作,血糖无法进入细胞,导致血糖升高。\n2. 疲劳:由于身体无法有效利用血糖产生能量,患者可能会感到疲劳无力。\n3. 肥胖:胰岛素抵抗可能导致过量的胰岛素分泌,使得脂肪存储增加,进而导致体重增加。\n4. 高血压:胰岛素抵抗与高血压之间存在关联,可能导致血压升高。\n5. 黑棘皮病:皮肤出现黑色、粗糙的皮疹,通常出现在颈部、腋窝等皮肤皱褶处。\n\n胰岛素抵抗的原因可能包括:\n\n1. 遗传因素:遗传易感基因可能会增加一个人患胰岛素抵抗的风险。\n2. 肥胖:过重或肥胖会增加患胰岛素抵抗的风险。\n3. 缺乏运动:缺乏运动可能会导致肌肉对胰岛素的敏感性降低。\n4. 高糖饮食:长期摄入高糖食品可能导致胰岛素抵抗。\n5. 睡眠不足:长期睡眠不足可能会影响胰岛素的正常功能。\n\n治疗胰岛素抵抗的方法主要包括生活方式的改变和药物治疗:\n\n1. 减轻体重:通过改善饮食和增加运动来减轻体重,有助于提高身体对胰岛素的敏感性。\n2. 增加运动:增加有氧运动和力量训练,有助于提高肌肉对胰岛素的敏感性。\n3. 改善饮食:摄入低糖、高纤维、富含全谷物和蔬菜水果的食物,有助于控制血糖和减轻体重。\n4. 药物治疗:在必要时,医生可能会开处降糖药物,如二甲双胍,以帮助提高胰岛素敏感性。\n\n胰岛素抵抗会影响一个人的日常生活,包括疲劳、体重增加、血糖波动等问题。通过改变生活习惯,如减轻体重、增加运动、改善饮食和保证充足睡眠,可以帮助改善胰岛素抵抗,降低患2型糖尿病的风险。"}]}
finetune/deepspeed_support/ds_zero2_no_offload.json ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "fp16": {
3
+ "enabled": "auto",
4
+ "loss_scale": 0,
5
+ "loss_scale_window": 1000,
6
+ "initial_scale_power": 16,
7
+ "hysteresis": 2,
8
+ "min_loss_scale": 1
9
+ },
10
+ "bf16": {
11
+ "enabled": "auto"
12
+ },
13
+ "zero_optimization": {
14
+ "stage": 2,
15
+ "allgather_partitions": true,
16
+ "allgather_bucket_size": 1e8,
17
+ "overlap_comm": true,
18
+ "reduce_scatter": true,
19
+ "reduce_bucket_size": 1e8,
20
+ "contiguous_gradients": true
21
+ },
22
+ "gradient_accumulation_steps": "auto",
23
+ "gradient_clipping": "auto",
24
+ "steps_per_print": 10,
25
+ "train_batch_size": "auto",
26
+ "train_micro_batch_size_per_gpu": "auto",
27
+ "wall_clock_breakdown": false
28
+ }
finetune/deepspeed_support/ds_zero2_offload.json ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "fp16": {
3
+ "enabled": "auto",
4
+ "loss_scale": 0,
5
+ "loss_scale_window": 1000,
6
+ "initial_scale_power": 16,
7
+ "hysteresis": 2,
8
+ "min_loss_scale": 1
9
+ },
10
+ "bf16": {
11
+ "enabled": "auto"
12
+ },
13
+
14
+ "zero_optimization": {
15
+ "stage": 2,
16
+ "offload_optimizer": {
17
+ "device": "cpu",
18
+ "pin_memory": true
19
+ },
20
+ "allgather_partitions": true,
21
+ "allgather_bucket_size": 5e8,
22
+ "overlap_comm": true,
23
+ "reduce_scatter": true,
24
+ "reduce_bucket_size": 5e8,
25
+ "contiguous_gradients": true
26
+ },
27
+
28
+ "gradient_accumulation_steps": "auto",
29
+ "gradient_clipping": "auto",
30
+ "steps_per_print": 10,
31
+ "train_batch_size": "auto",
32
+ "train_micro_batch_size_per_gpu": "auto",
33
+ "wall_clock_breakdown": false
34
+ }
finetune/deepspeed_support/ds_zero3_no_offload.json ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "fp16": {
3
+ "enabled": "auto",
4
+ "loss_scale": 0,
5
+ "loss_scale_window": 1000,
6
+ "initial_scale_power": 16,
7
+ "hysteresis": 2,
8
+ "min_loss_scale": 1
9
+ },
10
+ "bf16": {
11
+ "enabled": "auto"
12
+ },
13
+
14
+ "zero_optimization": {
15
+ "stage": 3,
16
+ "offload_optimizer": {
17
+ "device": "none",
18
+ "pin_memory": true
19
+ },
20
+ "offload_param": {
21
+ "device": "none",
22
+ "pin_memory": true
23
+ },
24
+ "overlap_comm": true,
25
+ "contiguous_gradients": true,
26
+ "sub_group_size": 1e9,
27
+ "reduce_bucket_size": "auto",
28
+ "stage3_prefetch_bucket_size": "auto",
29
+ "stage3_param_persistence_threshold": "auto",
30
+ "stage3_max_live_parameters": 1e9,
31
+ "stage3_max_reuse_distance": 1e9,
32
+ "stage3_gather_16bit_weights_on_model_save": false
33
+ },
34
+
35
+ "gradient_accumulation_steps": "auto",
36
+ "gradient_clipping": "auto",
37
+ "steps_per_print": 10,
38
+ "train_batch_size": "auto",
39
+ "train_micro_batch_size_per_gpu": "auto",
40
+ "wall_clock_breakdown": false
41
+ }
finetune/deepspeed_support/ds_zero3_offload.json ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "fp16": {
3
+ "enabled": "auto",
4
+ "loss_scale": 0,
5
+ "loss_scale_window": 1000,
6
+ "initial_scale_power": 16,
7
+ "hysteresis": 2,
8
+ "min_loss_scale": 1
9
+ },
10
+ "bf16": {
11
+ "enabled": "auto"
12
+ },
13
+
14
+ "zero_optimization": {
15
+ "stage": 3,
16
+ "offload_optimizer": {
17
+ "device": "cpu",
18
+ "pin_memory": true
19
+ },
20
+ "offload_param": {
21
+ "device": "cpu",
22
+ "pin_memory": true
23
+ },
24
+ "overlap_comm": true,
25
+ "contiguous_gradients": true,
26
+ "sub_group_size": 1e9,
27
+ "reduce_bucket_size": "auto",
28
+ "stage3_prefetch_bucket_size": "auto",
29
+ "stage3_param_persistence_threshold": "auto",
30
+ "stage3_max_live_parameters": 1e9,
31
+ "stage3_max_reuse_distance": 1e9,
32
+ "stage3_gather_16bit_weights_on_model_save": false
33
+ },
34
+
35
+ "gradient_accumulation_steps": "auto",
36
+ "gradient_clipping": "auto",
37
+ "steps_per_print": 10,
38
+ "train_batch_size": "auto",
39
+ "train_micro_batch_size_per_gpu": "auto",
40
+ "wall_clock_breakdown": false
41
+ }
finetune/deepspeed_support/train.py ADDED
@@ -0,0 +1,603 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2024 Tencent Inc. All Rights Reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ # Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.
16
+ #
17
+ # This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
18
+ # and OPT implementations in this library. It has been modified from its
19
+ # original forms to accommodate minor architectural differences compared
20
+ # to GPT-NeoX and OPT used by the Meta AI team that trained the model.
21
+ #
22
+ # Licensed under the Apache License, Version 2.0 (the "License");
23
+ # you may not use this file except in compliance with the License.
24
+ # You may obtain a copy of the License at
25
+ #
26
+ # http://www.apache.org/licenses/LICENSE-2.0
27
+ #
28
+ # Unless required by applicable law or agreed to in writing, software
29
+ # distributed under the License is distributed on an "AS IS" BASIS,
30
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
31
+ # See the License for the specific language governing permissions and
32
+ # limitations under the License.
33
+
34
+
35
+ import os
36
+ import sys
37
+ sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
38
+ import json
39
+ import torch
40
+ import shutil
41
+ import logging
42
+ from dataclasses import dataclass, field
43
+ from typing import Optional, Dict
44
+
45
+ import transformers
46
+ from torch.utils.data import Dataset
47
+ from transformers import Trainer, TrainerCallback
48
+ from peft import LoraConfig, get_peft_model, PeftModel
49
+ from transformers.trainer_utils import PREFIX_CHECKPOINT_DIR
50
+ from transformers.modeling_utils import unwrap_model
51
+
52
+
53
+ def print_args(args, name='arguments'):
54
+ """Print arguments."""
55
+ if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0:
56
+ print(f'------------------------ {name} ------------------------', flush=True)
57
+ str_list = []
58
+ for arg in vars(args):
59
+ dots = '.' * (48 - len(arg))
60
+ str_list.append(' {} {} {}'.format(arg, dots, getattr(args, arg)))
61
+ for arg in sorted(str_list, key=lambda x: x.lower()):
62
+ print(arg, flush=True)
63
+ print(f'-------------------- end of {name} ---------------------', flush=True)
64
+
65
+
66
+ @dataclass
67
+ class ModelArguments:
68
+ use_flash_attn: bool = field(
69
+ default=False,
70
+ metadata={"help": "Enable FlashAttention-2 for faster training."}
71
+ )
72
+ use_lora: bool = field(default=False, metadata={"help": "Enable Lora for faster training."})
73
+ lora_rank: int = field(default=64, metadata={"help": "The rank of lora."})
74
+ lora_alpha: int = field(default=8, metadata={"help": "Lora alpha"})
75
+ lora_dropout: float = field(default=0.0, metadata={"help": "Lora dropout"})
76
+
77
+
78
+ @dataclass
79
+ class DataArguments:
80
+ train_data_file: str = field(default=None, metadata={"help": "Path to the training data."})
81
+ max_seq_length: int = field(
82
+ default=2048,
83
+ metadata={"help": "The max sequence length of the model inputs after tokenization."}
84
+ )
85
+ complex_data: Optional[str] = field(default=None)
86
+ use_dummy_data: bool = field(default=False, metadata={"help": "Use dummy data."})
87
+
88
+
89
+ @dataclass
90
+ class TrainingArguments(transformers.TrainingArguments):
91
+ cache_dir: Optional[str] = field(default=None)
92
+ optim: str = field(default="adamw_torch")
93
+ tokenizer_name_or_path: Optional[str] = field(default=None)
94
+ model_name_or_path: Optional[str] = field(default=None)
95
+ min_lr: float = field(
96
+ default=0.01,
97
+ metadata={"help": "The final learning rate at the end of the decay will be learning_rate * min_lr"}
98
+ )
99
+
100
+
101
+ IGNORE_INDEX = -100
102
+
103
+ HY_START_ID = 120000
104
+ HY_MIDDLE_ID = 120001
105
+ HY_END_ID = 120025
106
+
107
+
108
+ class DummyDataset(Dataset):
109
+ def __init__(self, tokenizer, max_seq_length=512, length=1000):
110
+ self.tokenizer = tokenizer
111
+ self.max_seq_length = max_seq_length
112
+ self.length = length
113
+
114
+ def __len__(self):
115
+ return self.length
116
+
117
+ def __getitem__(self, index):
118
+ tokens = torch.randint(0, self.tokenizer.vocab_size, (self.max_seq_length, ))
119
+ return {'input_ids': tokens, 'labels': tokens}
120
+
121
+
122
+ class SFTDataset(Dataset):
123
+ def __init__(self, data_file, tokenizer, max_seq_length = 2048, prompt_format = 'mplus'):
124
+ self.tokenizer = tokenizer
125
+ self.prompt_format = prompt_format
126
+ self.max_seq_length = max_seq_length
127
+
128
+ self.data_list = self.load_data(data_file)
129
+
130
+ # Pre-compute special token IDs for loss masking (Scheme B)
131
+ self.hy_start_id = HY_START_ID
132
+ self.hy_middle_id = HY_MIDDLE_ID
133
+ self.hy_end_id = HY_END_ID
134
+ # "assistant" is encoded as BPE tokens [611, 10372] ("ass" + "istant")
135
+ self.assistant_bpe_ids = tokenizer.encode('assistant', add_special_tokens=False)
136
+ self.pad_token_id = tokenizer.pad_token_id
137
+
138
+ def __len__(self):
139
+ return len(self.data_list)
140
+
141
+ def load_data(self, data_file):
142
+ logging.info('Loading data: {}'.format(data_file))
143
+ with open(data_file, 'r', encoding='utf8') as f:
144
+ data_list = f.readlines()
145
+ logging.info("there are {} data in dataset".format(len(data_list)))
146
+ return data_list
147
+
148
+ def _find_assistant_turn_boundaries(self, token_ids):
149
+ """Find assistant turn boundaries using Scheme B.
150
+
151
+ Locates assistant turns by finding <hy_start> + "assistant" BPE tokens pattern,
152
+ then marks loss from the <hy_middle> token (exclusive) to <hy_end> (inclusive).
153
+
154
+ Returns:
155
+ List of (start, end) tuples where:
156
+ - start: position of <hy_middle> in the assistant turn (loss starts at start+1)
157
+ - end: position of <hy_end> for that turn (loss includes this position)
158
+ """
159
+ boundaries = []
160
+ assistant_len = len(self.assistant_bpe_ids)
161
+ ids_list = token_ids.tolist() if isinstance(token_ids, torch.Tensor) else token_ids
162
+ n = len(ids_list)
163
+
164
+ i = 0
165
+ while i < n:
166
+ # Look for <hy_start> token
167
+ if ids_list[i] == self.hy_start_id:
168
+ # Check if followed by "assistant" BPE tokens
169
+ role_start = i + 1
170
+ role_end = role_start + assistant_len
171
+ if role_end <= n and ids_list[role_start:role_end] == self.assistant_bpe_ids:
172
+ # Check if followed by <hy_middle>
173
+ if role_end < n and ids_list[role_end] == self.hy_middle_id:
174
+ middle_pos = role_end
175
+ # Find the corresponding <hy_end>
176
+ end_pos = None
177
+ for j in range(middle_pos + 1, n):
178
+ if ids_list[j] == self.hy_end_id:
179
+ end_pos = j
180
+ break
181
+ if end_pos is not None:
182
+ boundaries.append((middle_pos, end_pos))
183
+ i = end_pos + 1
184
+ continue
185
+ i += 1
186
+
187
+ return boundaries
188
+
189
+ def encode_data(self, data_dict):
190
+ model_inputs = {}
191
+ reasoning_effort = data_dict.get('reasoning_effort', None)
192
+ if reasoning_effort is None:
193
+ reasoning_effort = 'no_think'
194
+ try:
195
+ template_output = self.tokenizer.apply_chat_template(
196
+ data_dict['messages'], tokenize=True, return_dict=False,
197
+ reasoning_effort=reasoning_effort
198
+ )
199
+ except Exception as e:
200
+ print(f"[ERROR] apply_chat_template failed: {e}")
201
+ print(f"[ERROR] messages: {data_dict['messages']}")
202
+ print(f"[ERROR] reasoning_effort: {reasoning_effort}")
203
+ template_output = []
204
+
205
+ # Debug: Check template_output type and content
206
+ if isinstance(template_output, bool):
207
+ print(f"[WARNING] apply_chat_template returned bool: {template_output}")
208
+ print(f"[WARNING] messages: {data_dict['messages']}")
209
+ print(f"[WARNING] reasoning_effort: {reasoning_effort}")
210
+ template_output = []
211
+
212
+ if isinstance(template_output, list) and len(template_output) > 0 and isinstance(template_output[0], list):
213
+ template_output = template_output[0]
214
+
215
+ # Ensure template_output is a list of integers
216
+ if not isinstance(template_output, list) or not all(isinstance(x, int) for x in template_output):
217
+ print(f"[WARNING] Invalid template_output format: {type(template_output)}, content: {template_output}")
218
+ print(f"[WARNING] messages: {data_dict['messages']}")
219
+ template_output = []
220
+
221
+ message_tokens = torch.tensor(template_output, dtype=torch.long)
222
+
223
+ # Handle empty message_tokens case
224
+ if message_tokens.numel() == 0:
225
+ print(f"[WARNING] Empty message_tokens, skipping data sample")
226
+ input_ids = torch.tensor([], dtype=torch.long)
227
+ labels = torch.tensor([], dtype=torch.long)
228
+ attention_mask = torch.tensor([], dtype=torch.bool)
229
+ else:
230
+ # Scheme B: Find assistant turn boundaries and build labels
231
+ boundaries = self._find_assistant_turn_boundaries(message_tokens)
232
+ message_labels = torch.full_like(message_tokens, IGNORE_INDEX)
233
+
234
+ for middle_pos, end_pos in boundaries:
235
+ # Compute loss from the token after <hy_middle> to <hy_end> (inclusive)
236
+ message_labels[middle_pos + 1:end_pos + 1] = message_tokens[middle_pos + 1:end_pos + 1]
237
+
238
+ input_ids = message_tokens.to(torch.long)
239
+ labels = message_labels.to(torch.long)
240
+
241
+ input_ids = input_ids[:self.max_seq_length]
242
+ labels = labels[:self.max_seq_length]
243
+ attention_mask = input_ids.ne(self.pad_token_id).to(torch.bool)
244
+
245
+ model_inputs["input_ids"] = input_ids
246
+ model_inputs["attention_mask"] = attention_mask
247
+ model_inputs["labels"] = labels
248
+
249
+ return model_inputs
250
+
251
+ def __getitem__(self, index):
252
+ data = self.data_list[index]
253
+ data = json.loads(data)
254
+ model_inputs = self.encode_data(data)
255
+
256
+ # Check if the encoded data is empty (due to tokenization failure)
257
+ if model_inputs["input_ids"].numel() == 0:
258
+ # Return a valid placeholder sample to avoid crash
259
+ eos_token_id = self.hy_end_id
260
+ pad_token_id = self.pad_token_id
261
+
262
+ # Create a minimal valid sequence
263
+ placeholder_tokens = [self.hy_start_id, eos_token_id]
264
+ placeholder_tokens = placeholder_tokens[:self.max_seq_length]
265
+
266
+ input_ids = torch.tensor(placeholder_tokens, dtype=torch.long)
267
+ labels = torch.tensor([IGNORE_INDEX, IGNORE_INDEX], dtype=torch.long)[:self.max_seq_length]
268
+ attention_mask = torch.tensor([1, 1], dtype=torch.bool)[:self.max_seq_length]
269
+
270
+ # Pad to max_seq_length if needed
271
+ if len(placeholder_tokens) < self.max_seq_length:
272
+ padding_length = self.max_seq_length - len(placeholder_tokens)
273
+ input_ids = torch.cat([input_ids, torch.full((padding_length,), pad_token_id, dtype=torch.long)])
274
+ labels = torch.cat([labels, torch.full((padding_length,), IGNORE_INDEX, dtype=torch.long)])
275
+ attention_mask = torch.cat([attention_mask, torch.zeros(padding_length, dtype=torch.bool)])
276
+
277
+ model_inputs = {
278
+ "input_ids": input_ids,
279
+ "attention_mask": attention_mask,
280
+ "labels": labels
281
+ }
282
+
283
+ return model_inputs
284
+
285
+
286
+ @dataclass
287
+ class DataCollatorForSupervisedDataset(object):
288
+ """Collate examples for supervised fine-tuning."""
289
+
290
+ tokenizer: transformers.PreTrainedTokenizer
291
+
292
+ def __call__(self, instances):
293
+ input_ids = [instance['input_ids'] for instance in instances]
294
+ labels = [instance['labels'] for instance in instances]
295
+ pad_token_id = self.tokenizer.pad_token_id
296
+ input_ids = torch.nn.utils.rnn.pad_sequence(input_ids, batch_first=True, padding_value=pad_token_id)
297
+ labels = torch.nn.utils.rnn.pad_sequence(labels, batch_first=True, padding_value=IGNORE_INDEX)
298
+ return dict(
299
+ input_ids=input_ids,
300
+ labels=labels,
301
+ attention_mask=input_ids.ne(pad_token_id),
302
+ )
303
+
304
+
305
+ def make_supervised_data_module(tokenizer, data_args) -> Dict:
306
+ """Make dataset and collator for supervised fine-tuning."""
307
+ if data_args.use_dummy_data:
308
+ train_dataset = DummyDataset(tokenizer, data_args.max_seq_length)
309
+ else:
310
+ train_dataset = SFTDataset(
311
+ tokenizer=tokenizer,
312
+ data_file=data_args.train_data_file,
313
+ max_seq_length=data_args.max_seq_length
314
+ )
315
+ data_collator = DataCollatorForSupervisedDataset(tokenizer=tokenizer)
316
+ return dict(train_dataset=train_dataset, eval_dataset=None, data_collator=data_collator)
317
+
318
+
319
+ # Copy tokenizer and config files when saving checkpoints
320
+ class CustomSaveCallback(TrainerCallback):
321
+ def on_save(self, args, state, control, **kwargs):
322
+ if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0:
323
+ output_dir = os.path.join(args.output_dir, f"{PREFIX_CHECKPOINT_DIR}-{state.global_step}")
324
+
325
+ # Copy tokenizer and config files to checkpoint directory
326
+ tokenizer_files = [
327
+ 'config.json',
328
+ 'generation_config.json',
329
+ 'tokenizer_config.json',
330
+ 'tokenizer.json',
331
+ 'chat_template.jinja',
332
+ 'preprocessor_config.json',
333
+ 'hy.tiktoken',
334
+ 'tokenization_hy.py',
335
+ 'special_tokens_map.json',
336
+ ]
337
+ for fname in tokenizer_files:
338
+ src = os.path.join(args.tokenizer_name_or_path, fname)
339
+ if os.path.isfile(src):
340
+ shutil.copy(src, os.path.join(output_dir, fname))
341
+
342
+ return control
343
+
344
+
345
+ def train():
346
+ parser = transformers.HfArgumentParser((ModelArguments, DataArguments, TrainingArguments))
347
+ model_args, data_args, training_args = parser.parse_args_into_dataclasses()
348
+ print_args(model_args, 'model arguments')
349
+ print_args(data_args, 'data arguments')
350
+ print_args(training_args, 'training arguments')
351
+
352
+ tokenizer = transformers.AutoTokenizer.from_pretrained(
353
+ training_args.tokenizer_name_or_path,
354
+ trust_remote_code = True
355
+ )
356
+
357
+ init_kwargs = {}
358
+ if model_args.use_flash_attn:
359
+ init_kwargs["attn_implementation"] = "flash_attention_2"
360
+ # Workaround: transformers >= 5.x uses importlib.metadata.packages_distributions()
361
+ # to verify flash-attn package name, which fails when the package is installed under
362
+ # a custom distribution name (e.g. ptm-flash-attn). Patch the check to skip it.
363
+ try:
364
+ from transformers.modeling_flash_attention_utils import FLASH_ATTENTION_COMPATIBILITY_MATRIX
365
+ _orig_pkg_check = FLASH_ATTENTION_COMPATIBILITY_MATRIX[2]["pkg_availability_check"]
366
+ FLASH_ATTENTION_COMPATIBILITY_MATRIX[2]["pkg_availability_check"] = lambda *a, **kw: True
367
+ print("[Patch] Bypassed flash_attn package distribution name check for FA2.")
368
+ except Exception as e:
369
+ print(f"[Patch] Could not patch FA2 pkg check (non-fatal): {e}")
370
+
371
+ # Determine torch dtype
372
+ if training_args.bf16:
373
+ torch_dtype = torch.bfloat16
374
+ elif training_args.fp16:
375
+ torch_dtype = torch.float16
376
+ else:
377
+ torch_dtype = torch.float32
378
+
379
+ # -----------------------------------------------------------------------
380
+ # DeepSpeed ZeRO-3: Tell transformers that we are using ZeRO-3 so that
381
+ # from_pretrained will shard the model across ranks instead of loading
382
+ # the full model on each node's CPU (which would OOM for large models).
383
+ #
384
+ # NOTE: The ds_config may contain "auto" for batch size fields, which
385
+ # DeepSpeed cannot parse at this stage (before Trainer resolves them).
386
+ # We must fill in concrete values before passing to HfDeepSpeedConfig.
387
+ # -----------------------------------------------------------------------
388
+ if training_args.deepspeed:
389
+ from transformers.integrations import HfDeepSpeedConfig
390
+
391
+ # Load ds_config and resolve "auto" batch size fields
392
+ ds_config = training_args.deepspeed
393
+ if isinstance(ds_config, str):
394
+ with open(ds_config, 'r') as f:
395
+ ds_config = json.load(f)
396
+
397
+ # Fill in batch size fields that DeepSpeed needs for zero.Init
398
+ if ds_config.get("train_micro_batch_size_per_gpu", "auto") == "auto":
399
+ ds_config["train_micro_batch_size_per_gpu"] = training_args.per_device_train_batch_size
400
+ if ds_config.get("gradient_accumulation_steps", "auto") == "auto":
401
+ ds_config["gradient_accumulation_steps"] = training_args.gradient_accumulation_steps
402
+ if ds_config.get("train_batch_size", "auto") == "auto":
403
+ ds_config["train_batch_size"] = (
404
+ training_args.per_device_train_batch_size
405
+ * training_args.gradient_accumulation_steps
406
+ * (torch.distributed.get_world_size() if torch.distributed.is_initialized() else 1)
407
+ )
408
+
409
+ dschf = HfDeepSpeedConfig(ds_config) # noqa: F841 - must keep ref to avoid GC
410
+
411
+ # Check if model weights exist (not just the directory)
412
+ _has_weights = (
413
+ training_args.model_name_or_path is not None
414
+ and os.path.isdir(training_args.model_name_or_path)
415
+ and any(
416
+ os.path.isfile(os.path.join(training_args.model_name_or_path, f))
417
+ for f in ("model.safetensors", "pytorch_model.bin", "model.safetensors.index.json", "pytorch_model.bin.index.json")
418
+ )
419
+ )
420
+
421
+ if _has_weights:
422
+ print(f"Loading model from: {training_args.model_name_or_path}")
423
+ model = transformers.AutoModelForCausalLM.from_pretrained(
424
+ training_args.model_name_or_path,
425
+ trust_remote_code=True,
426
+ dtype=torch_dtype,
427
+ attn_implementation=init_kwargs.get("attn_implementation", None),
428
+ )
429
+ print(f"[HY4] Model loaded successfully via from_pretrained.")
430
+ else:
431
+ if training_args.model_name_or_path is None:
432
+ raise ValueError(
433
+ "--model_name_or_path must be specified. Cannot load model config from None. "
434
+ "Please provide the path to the model directory."
435
+ )
436
+ print(f"Model weights not found at: {training_args.model_name_or_path}, "
437
+ f"using random initialized model instead.")
438
+ config = transformers.AutoConfig.from_pretrained(
439
+ training_args.model_name_or_path,
440
+ trust_remote_code=True
441
+ )
442
+ model = transformers.AutoModelForCausalLM.from_config(
443
+ config,
444
+ trust_remote_code=True,
445
+ dtype=torch_dtype,
446
+ attn_implementation=init_kwargs.get("attn_implementation", None),
447
+ )
448
+
449
+ if model_args.use_lora:
450
+ # HY4 uses MLA (Multi-head Latent Attention) with different projection names
451
+ lora_config = LoraConfig(
452
+ r=model_args.lora_rank,
453
+ lora_alpha=model_args.lora_alpha,
454
+ lora_dropout=model_args.lora_dropout,
455
+ target_modules=["q_a_proj", "q_b_proj", "kv_a_proj_with_mqa", "kv_b_proj", "o_proj"],
456
+ bias="none",
457
+ task_type="CAUSAL_LM",
458
+ )
459
+ model = get_peft_model(model, lora_config)
460
+
461
+ # Fix: Mark PEFT LoRA wrapper modules as ZeRO-3 leaf modules.
462
+ # PEFT wraps target Linear layers with lora.Linear, adding extra
463
+ # sub-modules (base_layer, lora_A, lora_B). This changes the module
464
+ # tree structure and disrupts ZeRO-3's parameter fetch/release
465
+ # scheduling, causing OOM during backward recomputation.
466
+ # By marking these wrappers as z3_leaf, ZeRO-3 treats them as atomic
467
+ # units (same as the original Linear), restoring correct scheduling.
468
+ from deepspeed.utils import set_z3_leaf_module
469
+ from peft.tuners.lora import Linear as LoraLinear
470
+ z3_leaf_count = 0
471
+ for module in model.modules():
472
+ if isinstance(module, LoraLinear):
473
+ set_z3_leaf_module(module, True)
474
+ z3_leaf_count += 1
475
+ print(f"[z3_leaf] Marked {z3_leaf_count} LoraLinear modules with _z3_leaf=True", flush=True)
476
+
477
+ # Verify the attribute is actually set
478
+ verified_count = 0
479
+ for name, module in model.named_modules():
480
+ if isinstance(module, LoraLinear):
481
+ has_attr = getattr(module, '_z3_leaf', False)
482
+ if has_attr:
483
+ verified_count += 1
484
+ else:
485
+ print(f"[z3_leaf] WARNING: module '{name}' is LoraLinear but _z3_leaf={has_attr}", flush=True)
486
+ print(f"[z3_leaf] Verification after marking: {verified_count}/{z3_leaf_count} modules have _z3_leaf=True", flush=True)
487
+
488
+ data_module = make_supervised_data_module(tokenizer=tokenizer, data_args=data_args)
489
+ # Tell Trainer not to attempt DataParallel
490
+ model.is_parallelizable = True
491
+ model.model_parallel = True
492
+
493
+ training_args.lr_scheduler_kwargs = {
494
+ 'min_lr_rate': training_args.min_lr / training_args.learning_rate,
495
+ }
496
+
497
+ # -----------------------------------------------------------------------
498
+ # Fix: DeepSpeed ZeRO-3 + gradient checkpointing compatibility.
499
+ #
500
+ # PyTorch's torch.utils.checkpoint with use_reentrant=False (the default
501
+ # in transformers) performs strict metadata checks on recomputed tensors
502
+ # during backward. Under ZeRO-3, parameters are all-gathered during the
503
+ # first forward pass (shape=[full_size]) but may be partitioned back
504
+ # (shape=[0]) when the checkpoint recomputes, causing a CheckpointError.
505
+ #
506
+ # Setting use_reentrant=True avoids this strict metadata check.
507
+ # -----------------------------------------------------------------------
508
+ if training_args.gradient_checkpointing and training_args.deepspeed:
509
+ training_args.gradient_checkpointing_kwargs = {"use_reentrant": True}
510
+
511
+ trainer = Trainer(
512
+ model=model,
513
+ processing_class=tokenizer,
514
+ args=training_args,
515
+ callbacks=[CustomSaveCallback],
516
+ **data_module
517
+ )
518
+ model.config.use_cache = False
519
+
520
+ # -----------------------------------------------------------------------
521
+ # Monkey-patch: fix dtype mismatch in DeepSpeed ZeRO-3 linear wrapper.
522
+ #
523
+ # By this point the DeepSpeed engine has been initialised by the Trainer
524
+ # and torch.nn.functional.linear has been replaced with
525
+ # zero3_linear_wrap. That wrapper does NOT auto-align input/weight
526
+ # dtypes before the matmul, causing "expected mat1 and mat2 to have the
527
+ # same dtype" errors in mixed-precision paths (e.g. enable_lm_head_fp32
528
+ # casts input to fp32 but weight remains bf16 under ZeRO-3).
529
+ #
530
+ # We wrap F.linear HERE (after DeepSpeed init) so that:
531
+ # 1. We are sure to capture the already-replaced function.
532
+ # 2. The dtype cast happens *outside* the autograd.Function, so
533
+ # gradient-checkpointing recompute sees identical tensor metadata.
534
+ # -----------------------------------------------------------------------
535
+ import torch.nn.functional as _F
536
+ _orig_F_linear = _F.linear
537
+
538
+ def _dtype_safe_linear(input, weight, bias=None):
539
+ if input.dtype != weight.dtype:
540
+ input = input.to(weight.dtype)
541
+ return _orig_F_linear(input, weight, bias)
542
+
543
+ _F.linear = _dtype_safe_linear
544
+ # -----------------------------------------------------------------------
545
+
546
+ # -----------------------------------------------------------------------
547
+ # Monkey-patch: skip grad norm calculation when max_grad_norm == 0.
548
+ #
549
+ # Under ZeRO-3 + CPU offload, DeepSpeed's complete_grad_norm_calculation
550
+ # all-gathers every gradient on CPU and does an ALLREDUCE to compute the
551
+ # global L2 norm. For a large model this is extremely slow and triggers an
552
+ # NCCL ALLREDUCE timeout (NumelIn=1) at optimizer step, even when clipping
553
+ # is disabled via gradient_clipping=0.0 (which only skips the clip, not the
554
+ # computation).
555
+ #
556
+ # When the user explicitly sets --max_grad_norm 0 we fully skip the norm
557
+ # computation (no all-gather, no ALLREDUCE) by returning 0.0 early.
558
+ #
559
+ # NOTE: We patch the CLASS method (not instance) because trainer.deepspeed
560
+ # is None at this point — the DeepSpeed engine is created inside
561
+ # trainer.train(). By patching the class, any future engine instance will
562
+ # inherit the patched method.
563
+ # -----------------------------------------------------------------------
564
+ if getattr(training_args, "max_grad_norm", None) == 0:
565
+ import torch as _torch
566
+ from deepspeed.runtime.zero.stage3 import DeepSpeedZeroOptimizer_Stage3 as _ZeRO3Optimizer
567
+
568
+ def _skip_get_norm_groups(self):
569
+ return [_torch.tensor(0.0)]
570
+
571
+ _ZeRO3Optimizer._get_norm_groups = _skip_get_norm_groups
572
+ logging.info("[grad_norm] max_grad_norm=0: patched DeepSpeedZeroOptimizer_Stage3._get_norm_groups to skip norm computation")
573
+ # -----------------------------------------------------------------------
574
+
575
+ # -----------------------------------------------------------------------
576
+ # Post-DeepSpeed-init verification: check if _z3_leaf marks survived
577
+ # DeepSpeed engine initialization (which happens inside Trainer.__init__).
578
+ # -----------------------------------------------------------------------
579
+ if model_args.use_lora:
580
+ from peft.tuners.lora import Linear as LoraLinear
581
+ post_init_count = 0
582
+ post_init_verified = 0
583
+ for name, module in trainer.model.named_modules():
584
+ if isinstance(module, LoraLinear):
585
+ post_init_count += 1
586
+ has_attr = getattr(module, '_z3_leaf', False)
587
+ if has_attr:
588
+ post_init_verified += 1
589
+ elif post_init_count <= 5: # Only print first few warnings to avoid spam
590
+ print(f"[z3_leaf] POST-INIT WARNING: module '{name}' lost _z3_leaf after Trainer init!", flush=True)
591
+ print(f"[z3_leaf] Post-Trainer-init verification: {post_init_verified}/{post_init_count} LoraLinear modules still have _z3_leaf=True", flush=True)
592
+ # -----------------------------------------------------------------------
593
+
594
+ trainer.train(resume_from_checkpoint=training_args.resume_from_checkpoint)
595
+
596
+ # Synchronize all processes before exit to avoid "Connection reset by peer"
597
+ # warnings caused by timing differences in multi-node shutdown.
598
+ if torch.distributed.is_initialized():
599
+ torch.distributed.barrier()
600
+
601
+
602
+ if __name__ == "__main__":
603
+ train()
finetune/deepspeed_support/train.sh ADDED
@@ -0,0 +1,103 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/bin/bash
2
+
3
+ NET_TYPE="high"
4
+ export NCCL_DEBUG=WARN
5
+ export NCCL_P2P_LEVEL=NVL
6
+ export NCCL_IB_TIMEOUT=24
7
+ export NCCL_NVLS_ENABLE=0
8
+ export NCCL_MPI_PROFILE_PRIMS_ENABLE=0
9
+ export CUDA_DEVICE_MAX_CONNECTIONS=1
10
+ export TORCH_NCCL_HEARTBEAT_TIMEOUT_SEC=3600
11
+ if [[ "${NET_TYPE}" = "low" ]]; then
12
+ export NCCL_SOCKET_IFNAME=eth1
13
+ export NCCL_IB_GID_INDEX=3
14
+ export NCCL_IB_HCA=mlx5_2:1
15
+ export NCCL_IB_SL=3
16
+ export NCCL_CHECK_DISABLE=1
17
+ export NCCL_P2P_DISABLE=0
18
+ export NCCL_LL_THRESHOLD=16384
19
+ export NCCL_IB_CUDA_SUPPORT=1
20
+ else
21
+ export NCCL_IB_GID_INDEX=3
22
+ export NCCL_IB_SL=3
23
+ export NCCL_CHECK_DISABLE=1
24
+ export NCCL_P2P_DISABLE=0
25
+ export NCCL_IB_DISABLE=0
26
+ export NCCL_LL_THRESHOLD=16384
27
+ export NCCL_IB_CUDA_SUPPORT=1
28
+ export NCCL_SOCKET_IFNAME=bond1
29
+ export UCX_NET_DEVICES=bond1
30
+ export NCCL_IB_HCA=mlx5_bond_1,mlx5_bond_5,mlx5_bond_3,mlx5_bond_7,mlx5_bond_4,mlx5_bond_8,mlx5_bond_2,mlx5_bond_6
31
+ export NCCL_COLLNET_ENABLE=0
32
+ export SHARP_COLL_ENABLE_SAT=0
33
+ export NCCL_NET_GDR_LEVEL=2
34
+ export NCCL_IB_QPS_PER_CONNECTION=4
35
+ export NCCL_IB_TC=160
36
+ export NCCL_PXN_DISABLE=1
37
+ fi
38
+
39
+ export HOST_GPU_NUM=8
40
+ # IP list, comma separated. e.g. "192.168.1.1,192.168.1.2" or single node "192.168.1.1"
41
+ IP_LIST=${IP_LIST:-"127.0.0.1"}
42
+
43
+ IFS=',' read -ra IP_ARRAY <<< "$IP_LIST"
44
+ export NODES=${#IP_ARRAY[@]}
45
+ export LOCAL_IP=${IP_ARRAY[0]}
46
+ NODE_IP_LIST=""
47
+ for ip in "${IP_ARRAY[@]}"; do
48
+ if [ -n "$NODE_IP_LIST" ]; then
49
+ NODE_IP_LIST="${NODE_IP_LIST},"
50
+ fi
51
+ NODE_IP_LIST="${NODE_IP_LIST}${ip}:${HOST_GPU_NUM}"
52
+ done
53
+ export NODE_IP_LIST
54
+ export NODE_NUM=$((${NODES} * ${HOST_GPU_NUM}))
55
+
56
+ model_path=/path/to/hy_ckpt
57
+ tokenizer_path=/path/to/hy_ckpt
58
+ train_data_file=../data/example_data.jsonl
59
+
60
+ ds_config_file=ds_zero3_offload.json
61
+
62
+ output_path=./output
63
+
64
+ mkdir -p ${output_path}
65
+
66
+ current_time=$(date "+%Y.%m.%d-%H.%M.%S")
67
+ log_file=${output_path}/"log_${current_time}.txt"
68
+
69
+ echo $NODE_IP_LIST > env.txt 2>&1
70
+ sed "s/:/ slots=/g" env.txt | sed "s/,/\n/g" > "hostfile"
71
+ sed "s/:.//g" env.txt | sed "s/,/\n/g" > "pssh.hosts"
72
+ export CHIEF_IP=$LOCAL_IP
73
+
74
+ if [ ${NODES} -gt 1 ]; then
75
+ HOST_PATH=hostfile
76
+ DS_ARGS="--hostfile=${HOST_PATH} --master_addr ${CHIEF_IP}"
77
+ else
78
+ DS_ARGS=""
79
+ fi
80
+
81
+ echo "NODES: ${NODES}, LOCAL_IP: ${LOCAL_IP}, NODE_IP_LIST: ${NODE_IP_LIST}"
82
+
83
+ deepspeed ${DS_ARGS} \
84
+ train.py \
85
+ --do_train \
86
+ --model_name_or_path ${model_path} \
87
+ --tokenizer_name_or_path ${tokenizer_path} \
88
+ --train_data_file ${train_data_file} \
89
+ --deepspeed ${ds_config_file} \
90
+ --output_dir ${output_path} \
91
+ --per_device_train_batch_size 1 \
92
+ --gradient_accumulation_steps 1 \
93
+ --gradient_checkpointing \
94
+ --lr_scheduler_type cosine_with_min_lr \
95
+ --logging_steps 1 \
96
+ --save_steps 20 \
97
+ --learning_rate 1e-5 \
98
+ --min_lr 1e-6 \
99
+ --warmup_steps 0 \
100
+ --save_strategy steps \
101
+ --bf16 \
102
+ --max_grad_norm 0 \
103
+ --max_seq_length 1024 | tee ${log_file}
finetune/deepspeed_support/train_lora.sh ADDED
@@ -0,0 +1,109 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/bin/bash
2
+
3
+ NET_TYPE="high"
4
+ export NCCL_DEBUG=WARN
5
+ export NCCL_P2P_LEVEL=NVL
6
+ export NCCL_IB_TIMEOUT=24
7
+ export NCCL_NVLS_ENABLE=0
8
+ export NCCL_MPI_PROFILE_PRIMS_ENABLE=0
9
+ export CUDA_DEVICE_MAX_CONNECTIONS=1
10
+ export TORCH_NCCL_HEARTBEAT_TIMEOUT_SEC=3600
11
+ export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True
12
+ if [[ "${NET_TYPE}" = "low" ]]; then
13
+ export NCCL_SOCKET_IFNAME=eth1
14
+ export NCCL_IB_GID_INDEX=3
15
+ export NCCL_IB_HCA=mlx5_2:1
16
+ export NCCL_IB_SL=3
17
+ export NCCL_CHECK_DISABLE=1
18
+ export NCCL_P2P_DISABLE=0
19
+ export NCCL_LL_THRESHOLD=16384
20
+ export NCCL_IB_CUDA_SUPPORT=1
21
+ else
22
+ export NCCL_IB_GID_INDEX=3
23
+ export NCCL_IB_SL=3
24
+ export NCCL_CHECK_DISABLE=1
25
+ export NCCL_P2P_DISABLE=0
26
+ export NCCL_IB_DISABLE=0
27
+ export NCCL_LL_THRESHOLD=16384
28
+ export NCCL_IB_CUDA_SUPPORT=1
29
+ export NCCL_SOCKET_IFNAME=bond1
30
+ export UCX_NET_DEVICES=bond1
31
+ export NCCL_IB_HCA=mlx5_bond_1,mlx5_bond_5,mlx5_bond_3,mlx5_bond_7,mlx5_bond_4,mlx5_bond_8,mlx5_bond_2,mlx5_bond_6
32
+ export NCCL_COLLNET_ENABLE=0
33
+ export SHARP_COLL_ENABLE_SAT=0
34
+ export NCCL_NET_GDR_LEVEL=2
35
+ export NCCL_IB_QPS_PER_CONNECTION=4
36
+ export NCCL_IB_TC=160
37
+ export NCCL_PXN_DISABLE=1
38
+ fi
39
+
40
+ export HOST_GPU_NUM=8
41
+ # IP list, comma separated. e.g. "192.168.1.1,192.168.1.2" or single node "192.168.1.1"
42
+ IP_LIST=${IP_LIST:-"127.0.0.1"}
43
+
44
+ IFS=',' read -ra IP_ARRAY <<< "$IP_LIST"
45
+ export NODES=${#IP_ARRAY[@]}
46
+ export LOCAL_IP=${IP_ARRAY[0]}
47
+ NODE_IP_LIST=""
48
+ for ip in "${IP_ARRAY[@]}"; do
49
+ if [ -n "$NODE_IP_LIST" ]; then
50
+ NODE_IP_LIST="${NODE_IP_LIST},"
51
+ fi
52
+ NODE_IP_LIST="${NODE_IP_LIST}${ip}:${HOST_GPU_NUM}"
53
+ done
54
+ export NODE_IP_LIST
55
+ export NODE_NUM=$((${NODES} * ${HOST_GPU_NUM}))
56
+
57
+
58
+ model_path=/path/to/hy_ckpt
59
+ tokenizer_path=/path/to/hy_ckpt
60
+ train_data_file=../data/example_data.jsonl
61
+
62
+ ds_config_file=ds_zero3_offload.json
63
+
64
+ output_path=./output
65
+
66
+ mkdir -p ${output_path}
67
+
68
+ current_time=$(date "+%Y.%m.%d-%H.%M.%S")
69
+ log_file=${output_path}/"log_${current_time}.txt"
70
+
71
+ echo $NODE_IP_LIST > env.txt 2>&1
72
+ sed "s/:/ slots=/g" env.txt | sed "s/,/\n/g" > "hostfile"
73
+ sed "s/:.//g" env.txt | sed "s/,/\n/g" > "pssh.hosts"
74
+ export CHIEF_IP=$LOCAL_IP
75
+
76
+ if [ ${NODES} -gt 1 ]; then
77
+ HOST_PATH=hostfile
78
+ DS_ARGS="--hostfile=${HOST_PATH} --master_addr ${CHIEF_IP}"
79
+ else
80
+ DS_ARGS=""
81
+ fi
82
+
83
+ echo "NODES: ${NODES}, LOCAL_IP: ${LOCAL_IP}, NODE_IP_LIST: ${NODE_IP_LIST}"
84
+
85
+ deepspeed ${DS_ARGS} \
86
+ train.py \
87
+ --do_train \
88
+ --model_name_or_path ${model_path} \
89
+ --tokenizer_name_or_path ${tokenizer_path} \
90
+ --train_data_file ${train_data_file} \
91
+ --deepspeed ${ds_config_file} \
92
+ --output_dir ${output_path} \
93
+ --per_device_train_batch_size 1 \
94
+ --gradient_accumulation_steps 1 \
95
+ --gradient_checkpointing \
96
+ --lr_scheduler_type cosine_with_min_lr \
97
+ --logging_steps 1 \
98
+ --save_steps 20 \
99
+ --learning_rate 1e-5 \
100
+ --min_lr 1e-6 \
101
+ --warmup_steps 0 \
102
+ --save_strategy steps \
103
+ --bf16 \
104
+ --max_grad_norm 0 \
105
+ --use_lora \
106
+ --lora_rank 64 \
107
+ --lora_alpha 128 \
108
+ --lora_dropout 0.1 \
109
+ --max_seq_length 512 | tee ${log_file}
finetune/llama_factory_support/dataset_info.json ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "hy_v4_demo": {
3
+ "file_name": "../data/example_data.jsonl",
4
+ "formatting": "sharegpt",
5
+ "columns": {
6
+ "messages": "messages"
7
+ },
8
+ "tags": {
9
+ "role_tag": "role",
10
+ "content_tag": "content",
11
+ "user_tag": "user",
12
+ "assistant_tag": "assistant",
13
+ "system_tag": "system"
14
+ }
15
+ }
16
+ }
finetune/llama_factory_support/hy_v4_full_sft.yaml ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ### model
2
+ model_name_or_path: /path/to/hy_ckpt
3
+ trust_remote_code: true
4
+
5
+ ### method
6
+ stage: sft
7
+ do_train: true
8
+ finetuning_type: full
9
+
10
+ deepspeed: ../deepspeed_support/ds_zero3_offload.json
11
+
12
+ ### dataset
13
+ dataset_dir: .
14
+ dataset: hy_v4_demo
15
+ template: hy_v4
16
+ # NOTE: To train slow-thinking (CoT), include <think:opensource>...</think:opensource>
17
+ # tags in assistant content. Without think tags, only fast-thinking is trained.
18
+ cutoff_len: 256
19
+ max_samples: 1000
20
+ overwrite_cache: true
21
+
22
+ ### output
23
+ output_dir: saves/hy_v4/full/sft
24
+ logging_steps: 1
25
+ save_steps: 500
26
+ plot_loss: true
27
+ overwrite_output_dir: true
28
+ save_only_model: false
29
+ report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow]
30
+
31
+ ### train
32
+ per_device_train_batch_size: 1
33
+ gradient_accumulation_steps: 1
34
+ learning_rate: 1.0e-5
35
+ num_train_epochs: 3.0
36
+ lr_scheduler_type: cosine_with_min_lr
37
+ lr_scheduler_kwargs:
38
+ min_lr_rate: 0.1 # min_lr / learning_rate = 1e-6 / 1e-5 = 0.1
39
+ warmup_steps: 0
40
+ bf16: true
41
+ optim: adamw_torch
42
+ gradient_checkpointing: true
43
+ gradient_checkpointing_kwargs:
44
+ use_reentrant: true
45
+ ddp_timeout: 180000000
46
+ flash_attn: auto
47
+ max_grad_norm: 0.0
48
+ resume_from_checkpoint: null
49
+
50
+ ### eval
51
+ # eval_dataset: alpaca_en_demo
52
+ # val_size: 0.1
53
+ # per_device_eval_batch_size: 1
54
+ # eval_strategy: steps
55
+ # eval_steps: 500
finetune/llama_factory_support/hy_v4_lora_sft.yaml ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ### model
2
+ model_name_or_path: /path/to/hy_ckpt
3
+ trust_remote_code: true
4
+
5
+ ### method
6
+ stage: sft
7
+ do_train: true
8
+ finetuning_type: lora
9
+
10
+
11
+
12
+ # Using FSDP1 with full sharding strategy.
13
+ fsdp: true
14
+ fsdp_config:
15
+ version: 1
16
+ fsdp_sharding_strategy: FULL_SHARD
17
+ fsdp_backward_prefetch: BACKWARD_PRE
18
+ fsdp_state_dict_type: FULL_STATE_DICT
19
+ fsdp_use_orig_params: true
20
+ fsdp_cpu_ram_efficient_loading: true
21
+ fsdp_offload_params: false
22
+ fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP
23
+ fsdp_transformer_layer_cls_to_wrap: HYV4DecoderLayer
24
+
25
+ ### LoRA parameters
26
+ # HY4 uses MLA (Multi-head Latent Attention) with these projection names
27
+ lora_rank: 64
28
+ lora_alpha: 128
29
+ lora_dropout: 0.05
30
+ lora_target: q_a_proj,q_b_proj,kv_a_proj_with_mqa,kv_b_proj,o_proj
31
+
32
+ ### dataset
33
+ dataset_dir: .
34
+ dataset: hy_v4_demo
35
+ template: hy_v4
36
+ cutoff_len: 1024
37
+ max_samples: 1000
38
+ overwrite_cache: true
39
+
40
+ ### output
41
+ output_dir: saves/hy_v4/lora/sft
42
+ logging_steps: 1
43
+ save_steps: 500
44
+ plot_loss: true
45
+ overwrite_output_dir: true
46
+ save_only_model: false
47
+ report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow]
48
+
49
+ ### train
50
+ per_device_train_batch_size: 1
51
+ gradient_accumulation_steps: 1
52
+ learning_rate: 2.0e-4
53
+ num_train_epochs: 3.0
54
+ lr_scheduler_type: cosine_with_min_lr
55
+ lr_scheduler_kwargs:
56
+ min_lr_rate: 0.1 # min_lr / learning_rate = 2e-5 / 2e-4 = 0.1
57
+ warmup_steps: 0
58
+ bf16: true
59
+ gradient_checkpointing: true
60
+ gradient_checkpointing_kwargs:
61
+ use_reentrant: true
62
+ ddp_timeout: 180000000
63
+ flash_attn: auto
64
+ resume_from_checkpoint: null
65
+
66
+ ### eval
67
+ # eval_dataset: alpaca_en_demo
68
+ # val_size: 0.1
69
+ # per_device_eval_batch_size: 1
70
+ # eval_strategy: steps
71
+ # eval_steps: 500
finetune/llama_factory_support/hy_v4_patches.py ADDED
@@ -0,0 +1,768 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ This module applies all necessary runtime patches.
3
+
4
+
5
+ Patches applied:
6
+ 1. ZeRO-3 key rename + expert fuse + buffer loading
7
+ Handles both inner format (original) and outer format (pre-converted) checkpoints.
8
+ - Key renaming: mlp.router.gate -> mlp.gate, etc.
9
+ - Per-expert -> 3D fuse: experts.N.gate_proj -> experts.gate_up_proj
10
+ - Buffer loading: e_score_correction_bias (ZeRO-3 only handles parameters)
11
+ 2. Tokenizer file copy (CustomSaveCallback)
12
+ Ensures each checkpoint directory is self-contained for inference.
13
+ """
14
+
15
+ import os
16
+ import re
17
+ import logging
18
+ import shutil
19
+ from typing import Optional
20
+
21
+ import torch
22
+
23
+ logger = logging.getLogger(__name__)
24
+
25
+ # ============================================================================
26
+ # Patch 1: Key rename + expert fuse + buffer loading for ZeRO-3
27
+ #
28
+ # The checkpoint may be in either inner format (original) or outer format
29
+ # (pre-converted by convert_ckpt_to_outer.py). This patch handles both:
30
+ # - Key renaming: mlp.router.gate -> mlp.gate, etc.
31
+ # - Per-expert -> 3D fuse: experts.N.gate_proj -> experts.gate_up_proj
32
+ # - Buffer loading: e_score_correction_bias (ZeRO-3 only handles parameters)
33
+ #
34
+ # If the checkpoint is already in outer format, the rename/fuse logic is
35
+ # effectively a no-op (no matching keys to transform).
36
+ # ============================================================================
37
+
38
+ # Key renames: checkpoint inner format -> model format
39
+ _CKPT_KEY_RENAMES = [
40
+ ("mlp.router.gate.", "mlp.gate."),
41
+ ("mlp.expert_bias", "mlp.e_score_correction_bias"),
42
+ ("mlp.shared_mlp.", "mlp.shared_experts."),
43
+ # Also handle even older checkpoints that use mlp.gate.wg
44
+ ("mlp.gate.wg.", "mlp.gate."),
45
+ ]
46
+
47
+ # Regex to match per-expert keys in checkpoint
48
+ # e.g. model.layers.10.mlp.experts.5.gate_proj.weight
49
+ _EXPERT_KEY_RE = re.compile(
50
+ r"^(.*\.mlp\.experts\.)(\d+)\.(gate_proj|up_proj|down_proj)\.weight$"
51
+ )
52
+
53
+
54
+ def _apply_buffer_loading_patch():
55
+ """Patch the DeepSpeed ZeRO-3 state_dict loader to handle:
56
+ 1. Key renaming (inner checkpoint format -> model format)
57
+ 2. Per-expert -> 3D fused expert tensors
58
+ 3. Manual buffer loading (e_score_correction_bias etc.)
59
+
60
+ ZeRO-3's _load_state_dict_into_zero3_model only handles named_parameters.
61
+ Buffers like e_score_correction_bias must be loaded manually.
62
+ """
63
+ try:
64
+ from transformers.integrations.deepspeed import (
65
+ _load_state_dict_into_zero3_model as _orig_load_zero3,
66
+ )
67
+ import transformers.integrations.deepspeed as _ds_mod
68
+ import transformers.modeling_utils as _mu_mod
69
+ except ImportError:
70
+ logger.warning(
71
+ "Could not import transformers.integrations.deepspeed; "
72
+ "buffer loading patch NOT applied."
73
+ )
74
+ return
75
+
76
+ def _patched_load_zero3(model_to_load, state_dict, *args, **kwargs):
77
+ # Step 1: Key rename + per-expert collection
78
+ new_sd = {}
79
+ expert_groups = {} # prefix -> {expert_idx -> {proj_name -> tensor}}
80
+
81
+ for k, v in state_dict.items():
82
+ m = _EXPERT_KEY_RE.match(k)
83
+ if m:
84
+ # Per-expert key: collect for fusion
85
+ prefix = m.group(1)
86
+ expert_idx = int(m.group(2))
87
+ proj_name = m.group(3)
88
+ if prefix not in expert_groups:
89
+ expert_groups[prefix] = {}
90
+ if expert_idx not in expert_groups[prefix]:
91
+ expert_groups[prefix][expert_idx] = {}
92
+ expert_groups[prefix][expert_idx][proj_name] = v
93
+ else:
94
+ # Non-expert key: apply simple renames
95
+ new_k = k
96
+ for old_sub, new_sub in _CKPT_KEY_RENAMES:
97
+ if old_sub in new_k:
98
+ new_k = new_k.replace(old_sub, new_sub)
99
+ new_sd[new_k] = v
100
+
101
+ # Step 2: Fuse expert groups into 3D tensors
102
+ if expert_groups:
103
+ for prefix in sorted(expert_groups.keys()):
104
+ experts = expert_groups[prefix]
105
+ num_experts = max(experts.keys()) + 1
106
+ gate_up_list = []
107
+ down_list = []
108
+ for i in range(num_experts):
109
+ if i not in experts:
110
+ logger.warning(
111
+ "HYV4 Patch 1: Missing expert %d in %s", i, prefix
112
+ )
113
+ continue
114
+ exp = experts[i]
115
+ if "gate_proj" in exp and "up_proj" in exp:
116
+ gate_up_list.append(
117
+ torch.cat([exp["gate_proj"], exp["up_proj"]], dim=0)
118
+ )
119
+ if "down_proj" in exp:
120
+ down_list.append(exp["down_proj"])
121
+ if gate_up_list:
122
+ new_sd[f"{prefix}gate_up_proj"] = torch.stack(gate_up_list, dim=0)
123
+ if down_list:
124
+ new_sd[f"{prefix}down_proj"] = torch.stack(down_list, dim=0)
125
+ logger.info(
126
+ "HYV4 Patch 1: Fused %d expert groups from per-expert to 3D format.",
127
+ len(expert_groups)
128
+ )
129
+ del expert_groups
130
+
131
+ # Step 3: Load parameters via original ZeRO-3 loader
132
+ result = _orig_load_zero3(model_to_load, new_sd, *args, **kwargs)
133
+
134
+ # Step 4: Manually load buffers (e.g. e_score_correction_bias)
135
+ # ZeRO-3's loader only handles named_parameters, not named_buffers.
136
+ buffers_loaded = 0
137
+ for name, buf in model_to_load.named_buffers():
138
+ if name in new_sd:
139
+ src_tensor = new_sd[name]
140
+ if isinstance(src_tensor, torch.Tensor):
141
+ buf.data.copy_(src_tensor.to(buf.dtype))
142
+ buffers_loaded += 1
143
+ if isinstance(result, tuple) and len(result) >= 2:
144
+ if isinstance(result[1], set):
145
+ result[1].discard(name)
146
+ if buffers_loaded > 0:
147
+ logger.info(
148
+ "HYV4 Patch 1: Manually loaded %d buffers into model.",
149
+ buffers_loaded
150
+ )
151
+
152
+ del new_sd
153
+ return result
154
+
155
+ _ds_mod._load_state_dict_into_zero3_model = _patched_load_zero3
156
+ _mu_mod._load_state_dict_into_zero3_model = _patched_load_zero3
157
+ logger.info(
158
+ "HYV4 patch applied: ZeRO-3 key rename + expert fuse + buffer loading."
159
+ )
160
+
161
+ # ============================================================================
162
+ # Patch 2: Tokenizer file copy callback
163
+ #
164
+ # Ensures each checkpoint directory is self-contained for inference by
165
+ # copying all tokenizer-related files from the original tokenizer path.
166
+ # ============================================================================
167
+
168
+ # Tokenizer files that should be copied to each checkpoint
169
+ _TOKENIZER_FILES = [
170
+ "config.json",
171
+ "generation_config.json",
172
+ "tokenizer_config.json",
173
+ "tokenizer.json",
174
+ "chat_template.jinja",
175
+ "preprocessor_config.json",
176
+ "hy.tiktoken",
177
+ "tokenization_hy.py",
178
+ "special_tokens_map.json",
179
+ ]
180
+
181
+ def _copy_tokenizer_to_checkpoint(tokenizer_dir: str, checkpoint_dir: str):
182
+ """Copy tokenizer files from tokenizer_dir to checkpoint_dir."""
183
+ for fname in _TOKENIZER_FILES:
184
+ src = os.path.join(tokenizer_dir, fname)
185
+ if os.path.isfile(src):
186
+ shutil.copy(src, os.path.join(checkpoint_dir, fname))
187
+
188
+ # ============================================================================
189
+ # LLaMA Factory Callback: integrates patch 2 (tokenizer copy) into the
190
+ # training loop.
191
+ # ============================================================================
192
+
193
+ try:
194
+ from transformers import TrainerCallback
195
+ from transformers.trainer_utils import PREFIX_CHECKPOINT_DIR
196
+
197
+ class HYV4PatchCallback(TrainerCallback):
198
+ """
199
+ LLaMA Factory compatible callback that applies HYV4-specific patches.
200
+
201
+ Add to your YAML or pass to Trainer:
202
+ callbacks: [hy_v4_patches.HYV4PatchCallback]
203
+ """
204
+
205
+ def __init__(self, tokenizer_dir: Optional[str] = None):
206
+ """
207
+ Args:
208
+ tokenizer_dir: Path to the original tokenizer directory.
209
+ If None, will try to use model_name_or_path from training args.
210
+ """
211
+ self._tokenizer_dir = tokenizer_dir
212
+
213
+ def on_save(self, args, state, control, **kwargs):
214
+ # --- Patch 2: Copy tokenizer files ---
215
+ if torch.distributed.is_initialized() and torch.distributed.get_rank() != 0:
216
+ return control
217
+
218
+ checkpoint_dir = os.path.join(
219
+ args.output_dir,
220
+ f"{PREFIX_CHECKPOINT_DIR}-{state.global_step}",
221
+ )
222
+
223
+ # Determine tokenizer directory
224
+ tokenizer_dir = self._tokenizer_dir
225
+ if tokenizer_dir is None:
226
+ # Try common locations
227
+ tokenizer_dir = getattr(args, "tokenizer_name_or_path", None)
228
+ if tokenizer_dir is None:
229
+ tokenizer_dir = getattr(args, "model_name_or_path", None)
230
+
231
+ if tokenizer_dir and os.path.isdir(tokenizer_dir):
232
+ _copy_tokenizer_to_checkpoint(tokenizer_dir, checkpoint_dir)
233
+ logger.info(
234
+ "HYV4: Copied tokenizer files from %s to %s",
235
+ tokenizer_dir, checkpoint_dir
236
+ )
237
+
238
+ return control
239
+
240
+ except ImportError:
241
+ logger.warning(
242
+ "transformers not available; HYV4PatchCallback not defined."
243
+ )
244
+
245
+ # ============================================================================
246
+ # Patch 3: Memory-efficient shard-by-shard model loading for ZeRO-3
247
+ #
248
+ # The default transformers from_pretrained + ZeRO-3 path loads ALL shards
249
+ # into a single state_dict in CPU memory before distributing. For a ~670GB
250
+ # model with 8 processes per node, this causes CPU OOM.
251
+ #
252
+ # This patch replaces from_pretrained with a shard-by-shard loader that:
253
+ # 1. Creates the model skeleton under deepspeed.zero.Init (meta tensors)
254
+ # 2. Loads each safetensors shard one at a time (~7GB each)
255
+ # 3. Applies key renames + expert fusion per shard
256
+ # 4. Scatters into ZeRO-3 partitions immediately
257
+ # 5. Frees the shard before loading the next one
258
+ #
259
+ # This reduces per-rank CPU memory from ~670GB to ~7GB.
260
+ # ============================================================================
261
+
262
+ def _apply_shard_loading_patch():
263
+ """Monkey-patch AutoModelForCausalLM.from_pretrained to use shard-by-shard
264
+ loading when DeepSpeed ZeRO-3 is active."""
265
+ import gc
266
+ import json as _json
267
+ import transformers
268
+
269
+ _orig_from_pretrained = transformers.AutoModelForCausalLM.from_pretrained
270
+
271
+ def _shard_loading_from_pretrained(pretrained_model_name_or_path, *args, **kwargs):
272
+ """Memory-efficient from_pretrained that loads shards one at a time."""
273
+
274
+ model_path = pretrained_model_name_or_path
275
+
276
+ # Helper: fallback to default from_pretrained with CPU-safe loading.
277
+ # When not using DeepSpeed ZeRO-3 (e.g. FSDP1 mode), we must avoid
278
+ # loading the full model onto GPU (which would OOM for large models).
279
+ # Force device_map to CPU; FSDP1 will handle GPU sharding after wrap.
280
+ def _fallback_load():
281
+ kwargs.setdefault("low_cpu_mem_usage", True)
282
+ # For FSDP mode: only local_rank 0 loads real weights to CPU;
283
+ # other ranks create model on meta device (zero CPU memory).
284
+ # FSDP's sync_module_states=True (default) will broadcast weights
285
+ # from rank 0 to all other ranks during wrap.
286
+ # This reduces per-node peak CPU memory from N_ranks * model_size
287
+ # to just 1 * model_size.
288
+ local_rank = int(os.environ.get("LOCAL_RANK", "0"))
289
+ if local_rank != 0:
290
+ logger.info(
291
+ "[HYV4] FSDP mode: local_rank=%d != 0, "
292
+ "creating model on meta device (zero CPU memory). "
293
+ "Weights will be synced from rank 0 by FSDP.",
294
+ local_rank,
295
+ )
296
+ import transformers
297
+ config = transformers.AutoConfig.from_pretrained(
298
+ pretrained_model_name_or_path,
299
+ trust_remote_code=kwargs.get("trust_remote_code", False),
300
+ )
301
+ torch_dtype = kwargs.get("torch_dtype", None)
302
+ if torch_dtype is None:
303
+ torch_dtype = getattr(config, "torch_dtype", torch.bfloat16)
304
+ if not isinstance(torch_dtype, torch.dtype):
305
+ torch_dtype = torch.bfloat16
306
+ with torch.device("meta"):
307
+ model = transformers.AutoModelForCausalLM.from_config(
308
+ config,
309
+ torch_dtype=torch_dtype,
310
+ trust_remote_code=kwargs.get("trust_remote_code", False),
311
+ )
312
+ return model
313
+
314
+ # local_rank 0: load real weights to CPU.
315
+ # FSDP1 will handle GPU sharding after wrap.
316
+ logger.info(
317
+ "[HYV4] FSDP mode: local_rank=0, loading real weights to CPU."
318
+ )
319
+ # LLaMA Factory may pass device_map pointing to GPU, which
320
+ # would cause OOM for large models.
321
+ kwargs["device_map"] = {"": "cpu"}
322
+ return _orig_from_pretrained(
323
+ pretrained_model_name_or_path, *args, **kwargs
324
+ )
325
+
326
+ # Only apply shard loading if:
327
+ # 1. It's a local directory with safetensors
328
+ # 2. DeepSpeed ZeRO-3 is being used
329
+ if not (isinstance(model_path, str) and os.path.isdir(model_path)):
330
+ return _fallback_load()
331
+
332
+ index_file = os.path.join(model_path, "model.safetensors.index.json")
333
+ single_file = os.path.join(model_path, "model.safetensors")
334
+ if not (os.path.isfile(index_file) or os.path.isfile(single_file)):
335
+ return _fallback_load()
336
+
337
+ # Try to determine the DeepSpeed config
338
+ ds_config = None
339
+
340
+ # Check if there's a deepspeed config in the HfTrainerDeepSpeedConfig
341
+ try:
342
+ from transformers.integrations.deepspeed import is_deepspeed_zero3_enabled
343
+ if not is_deepspeed_zero3_enabled():
344
+ logger.info(
345
+ "[HYV4 Patch 3] ZeRO-3 not enabled, using CPU fallback loader."
346
+ )
347
+ return _fallback_load()
348
+ except (ImportError, Exception):
349
+ # If we can't determine, try to proceed anyway
350
+ pass
351
+
352
+ # Get the deepspeed config from HF's global state
353
+ try:
354
+ from transformers.integrations.deepspeed import deepspeed_config as _get_ds_config
355
+ ds_config = _get_ds_config()
356
+ except (ImportError, Exception):
357
+ ds_config = None
358
+
359
+ if ds_config is None:
360
+ # Fallback: try weak ref approach
361
+ try:
362
+ from transformers.integrations import deepspeed as _hf_ds
363
+ if hasattr(_hf_ds, '_hf_deepspeed_config_weak_ref'):
364
+ _weak_ref = _hf_ds._hf_deepspeed_config_weak_ref
365
+ if _weak_ref is not None:
366
+ ds_obj = _weak_ref()
367
+ if ds_obj is not None:
368
+ ds_config = ds_obj.config
369
+ except (ImportError, AttributeError, Exception):
370
+ pass
371
+
372
+ if ds_config is None:
373
+ # Last resort: look for the config file path in environment
374
+ ds_config_path = os.environ.get("DEEPSPEED_CONFIG_FILE", None)
375
+ if ds_config_path is None:
376
+ ds_config_path = os.environ.get("DEEPSPEED_CONFIG", None)
377
+ if ds_config_path and os.path.isfile(ds_config_path):
378
+ with open(ds_config_path, "r") as f:
379
+ ds_config = _json.load(f)
380
+
381
+ if ds_config is None:
382
+ logger.warning(
383
+ "[HYV4 Patch 3] Cannot determine DeepSpeed config, "
384
+ "falling back to CPU loader (FSDP mode)."
385
+ )
386
+ return _fallback_load()
387
+
388
+ # Ensure ds_config is a dict
389
+ if hasattr(ds_config, 'config'):
390
+ ds_config = ds_config.config
391
+ if not isinstance(ds_config, dict):
392
+ logger.warning(
393
+ "[HYV4 Patch 3] ds_config is not a dict (%s), falling back.",
394
+ type(ds_config)
395
+ )
396
+ return _fallback_load()
397
+
398
+ # Check if it's actually ZeRO stage 3
399
+ zero_stage = ds_config.get("zero_optimization", {}).get("stage", 0)
400
+ if zero_stage != 3:
401
+ logger.info(
402
+ "[HYV4 Patch 3] Not ZeRO-3 (stage=%d), using CPU fallback loader.",
403
+ zero_stage
404
+ )
405
+ return _fallback_load()
406
+
407
+ logger.info(
408
+ "[HYV4 Patch 3] Using shard-by-shard loading for model at: %s",
409
+ model_path
410
+ )
411
+
412
+ import deepspeed
413
+
414
+ try:
415
+ from safetensors import safe_open
416
+ from transformers.integrations.deepspeed import (
417
+ _load_state_dict_into_zero3_model as _load_zero3,
418
+ )
419
+ except ImportError as e:
420
+ logger.warning(
421
+ "[HYV4 Patch 3] Required imports not available (%s), "
422
+ "falling back to default from_pretrained.", e
423
+ )
424
+ return _orig_from_pretrained(pretrained_model_name_or_path, *args, **kwargs)
425
+
426
+ # Replace "auto" values that deepspeed.zero.Init cannot resolve
427
+ ds_config_copy = _json.loads(_json.dumps(ds_config))
428
+ _auto_defaults = {
429
+ "train_batch_size": 32,
430
+ "train_micro_batch_size_per_gpu": 1,
431
+ "gradient_accumulation_steps": 1,
432
+ "gradient_clipping": 1.0,
433
+ }
434
+ for k, v in _auto_defaults.items():
435
+ if k in ds_config_copy and ds_config_copy[k] == "auto":
436
+ ds_config_copy[k] = v
437
+
438
+ # Determine dtype
439
+ torch_dtype = kwargs.pop("torch_dtype", torch.bfloat16)
440
+ if torch_dtype is None or torch_dtype == "auto":
441
+ torch_dtype = torch.bfloat16
442
+ trust_remote_code = kwargs.pop("trust_remote_code", True)
443
+ attn_implementation = kwargs.pop("attn_implementation", None)
444
+ # Pop config if already provided by caller (e.g. LLaMA Factory)
445
+ config = kwargs.pop("config", None)
446
+
447
+ # Step 1: Create model skeleton under ZeRO-3 Init (meta tensors)
448
+ if config is None:
449
+ config = transformers.AutoConfig.from_pretrained(
450
+ model_path, trust_remote_code=trust_remote_code
451
+ )
452
+ with deepspeed.zero.Init(
453
+ dtype=torch_dtype, config_dict_or_path=ds_config_copy
454
+ ):
455
+ model = transformers.AutoModelForCausalLM.from_config(
456
+ config,
457
+ trust_remote_code=trust_remote_code,
458
+ torch_dtype=torch_dtype,
459
+ attn_implementation=attn_implementation,
460
+ )
461
+ logger.info("[HYV4 Patch 3] Model skeleton created under ZeRO-3 Init.")
462
+
463
+ # Step 2: Determine shard files
464
+ if os.path.isfile(index_file):
465
+ with open(index_file, "r") as f:
466
+ index_data = _json.load(f)
467
+ shard_files = list(dict.fromkeys(index_data["weight_map"].values()))
468
+ else:
469
+ shard_files = ["model.safetensors"]
470
+
471
+ # Step 3: Load each shard and scatter into ZeRO-3 model
472
+ total_shards = len(shard_files)
473
+ all_loaded_keys = set()
474
+ pending_experts = {} # prefix -> {expert_idx -> {proj_name -> tensor}}
475
+
476
+ for shard_idx, shard_name in enumerate(shard_files, 1):
477
+ shard_path = os.path.join(model_path, shard_name)
478
+ logger.info(
479
+ "[HYV4 Patch 3] Loading shard %d/%d: %s",
480
+ shard_idx, total_shards, shard_name
481
+ )
482
+
483
+ # Load shard into CPU memory
484
+ shard_sd = {}
485
+ with safe_open(shard_path, framework="pt", device="cpu") as f:
486
+ for key in f.keys():
487
+ shard_sd[key] = f.get_tensor(key)
488
+
489
+ # Separate expert keys from non-expert keys, apply renames
490
+ renamed_sd = {}
491
+ expert_keys_in_shard = {}
492
+
493
+ for k, v in shard_sd.items():
494
+ m = _EXPERT_KEY_RE.match(k)
495
+ if m:
496
+ prefix = m.group(1)
497
+ expert_idx = int(m.group(2))
498
+ proj_name = m.group(3)
499
+ if prefix not in expert_keys_in_shard:
500
+ expert_keys_in_shard[prefix] = {}
501
+ if expert_idx not in expert_keys_in_shard[prefix]:
502
+ expert_keys_in_shard[prefix][expert_idx] = {}
503
+ expert_keys_in_shard[prefix][expert_idx][proj_name] = v
504
+ else:
505
+ new_k = k
506
+ for old_sub, new_sub in _CKPT_KEY_RENAMES:
507
+ if old_sub in new_k:
508
+ new_k = new_k.replace(old_sub, new_sub)
509
+ renamed_sd[new_k] = v
510
+ del shard_sd
511
+
512
+ # Merge expert keys into pending_experts
513
+ for prefix, experts in expert_keys_in_shard.items():
514
+ if prefix not in pending_experts:
515
+ pending_experts[prefix] = {}
516
+ for idx, projs in experts.items():
517
+ if idx not in pending_experts[prefix]:
518
+ pending_experts[prefix][idx] = {}
519
+ pending_experts[prefix][idx].update(projs)
520
+ del expert_keys_in_shard
521
+
522
+ # Check for completed expert groups
523
+ completed_prefixes = []
524
+ for prefix, experts in pending_experts.items():
525
+ if not experts:
526
+ continue
527
+ max_idx = max(experts.keys())
528
+ num_experts_found = len(experts)
529
+ all_complete = all(
530
+ len(projs) == 3 for projs in experts.values()
531
+ )
532
+ if all_complete and num_experts_found == (max_idx + 1):
533
+ completed_prefixes.append(prefix)
534
+
535
+ # Fuse completed expert groups
536
+ for prefix in completed_prefixes:
537
+ experts = pending_experts.pop(prefix)
538
+ num_experts_layer = max(experts.keys()) + 1
539
+ gate_up_list = []
540
+ down_list = []
541
+ for i in range(num_experts_layer):
542
+ exp = experts[i]
543
+ gate_up = torch.cat([exp["gate_proj"], exp["up_proj"]], dim=0)
544
+ gate_up_list.append(gate_up)
545
+ down_list.append(exp["down_proj"])
546
+ fused_gate_up = torch.stack(gate_up_list, dim=0)
547
+ fused_down = torch.stack(down_list, dim=0)
548
+ del gate_up_list, down_list, experts
549
+ renamed_sd[f"{prefix}gate_up_proj"] = fused_gate_up
550
+ renamed_sd[f"{prefix}down_proj"] = fused_down
551
+ logger.info(
552
+ "[HYV4 Patch 3] Fused %d experts for %s",
553
+ num_experts_layer, prefix
554
+ )
555
+
556
+ # Scatter this shard's weights into ZeRO-3 model
557
+ if renamed_sd:
558
+ _load_zero3(model, renamed_sd)
559
+ # Also load buffers
560
+ for name, buf in model.named_buffers():
561
+ if name in renamed_sd:
562
+ src_tensor = renamed_sd[name]
563
+ if isinstance(src_tensor, torch.Tensor):
564
+ buf.data.copy_(src_tensor.to(buf.dtype))
565
+ all_loaded_keys.update(renamed_sd.keys())
566
+ del renamed_sd
567
+ gc.collect()
568
+
569
+ # Flush remaining pending experts
570
+ if pending_experts:
571
+ logger.info(
572
+ "[HYV4 Patch 3] Flushing %d remaining expert group(s)...",
573
+ len(pending_experts)
574
+ )
575
+ flush_sd = {}
576
+ for prefix, experts in pending_experts.items():
577
+ num_experts_layer = max(experts.keys()) + 1
578
+ gate_up_list = []
579
+ down_list = []
580
+ for i in range(num_experts_layer):
581
+ if i not in experts:
582
+ logger.warning(
583
+ "[HYV4 Patch 3] Missing expert %d in %s", i, prefix
584
+ )
585
+ continue
586
+ exp = experts[i]
587
+ gate_up = torch.cat([exp["gate_proj"], exp["up_proj"]], dim=0)
588
+ gate_up_list.append(gate_up)
589
+ down_list.append(exp["down_proj"])
590
+ if gate_up_list:
591
+ fused_gate_up = torch.stack(gate_up_list, dim=0)
592
+ fused_down = torch.stack(down_list, dim=0)
593
+ flush_sd[f"{prefix}gate_up_proj"] = fused_gate_up
594
+ flush_sd[f"{prefix}down_proj"] = fused_down
595
+ logger.info(
596
+ "[HYV4 Patch 3] Fused %d experts for %s",
597
+ len(gate_up_list), prefix
598
+ )
599
+ del gate_up_list, down_list
600
+ del pending_experts
601
+
602
+ if flush_sd:
603
+ _load_zero3(model, flush_sd)
604
+ for name, buf in model.named_buffers():
605
+ if name in flush_sd:
606
+ src_tensor = flush_sd[name]
607
+ if isinstance(src_tensor, torch.Tensor):
608
+ buf.data.copy_(src_tensor.to(buf.dtype))
609
+ all_loaded_keys.update(flush_sd.keys())
610
+ del flush_sd
611
+ gc.collect()
612
+
613
+ # Report missing/unexpected keys
614
+ model_keys = set(n for n, _ in model.named_parameters())
615
+ model_keys.update(n for n, _ in model.named_buffers())
616
+ missing = model_keys - all_loaded_keys
617
+ unexpected = all_loaded_keys - model_keys
618
+ if missing:
619
+ real_missing = {k for k in missing if "lm_head" not in k}
620
+ if real_missing:
621
+ logger.warning(
622
+ "[HYV4 Patch 3] %d keys not found in checkpoint (first 10): %s",
623
+ len(real_missing), list(real_missing)[:10]
624
+ )
625
+ if unexpected:
626
+ logger.warning(
627
+ "[HYV4 Patch 3] %d unexpected keys (first 10): %s",
628
+ len(unexpected), list(unexpected)[:10]
629
+ )
630
+ logger.info(
631
+ "[HYV4 Patch 3] Shard-by-shard loading complete. "
632
+ "Loaded %d keys from %d shards.",
633
+ len(all_loaded_keys), total_shards
634
+ )
635
+
636
+ return model
637
+
638
+ # Apply the monkey-patch
639
+ transformers.AutoModelForCausalLM.from_pretrained = staticmethod(_shard_loading_from_pretrained)
640
+ logger.info(
641
+ "HYV4 patch applied: shard-by-shard model loading for ZeRO-3 "
642
+ "(reduces CPU memory from ~670GB to ~7GB per rank)."
643
+ )
644
+
645
+
646
+ # ============================================================================
647
+ # Auto-apply patches on import
648
+ # ============================================================================
649
+
650
+ # Patch 1: ZeRO-3 key rename + expert fuse + buffer loading
651
+ _apply_buffer_loading_patch()
652
+
653
+ # Patch 3: Memory-efficient shard-by-shard loading
654
+ _apply_shard_loading_patch()
655
+
656
+ # Patch 4: Unify model dtype before FSDP wrap
657
+ # FSDP1 requires all parameters in the same FSDP unit to have uniform dtype.
658
+ # After LoRA injection, adapter params are float32 while base model is bfloat16.
659
+ # This patch casts all parameters to bf16 before FSDP wraps the model.
660
+ def _apply_fsdp_dtype_patch():
661
+ """Monkey-patch Trainer to cast model to bf16 before FSDP wrap."""
662
+ try:
663
+ from transformers import Trainer
664
+
665
+ _orig_prepare_for_training = Trainer._prepare_for_training
666
+
667
+ def _patched_prepare_for_training(self, *args, **kwargs):
668
+ # Cast all model parameters to bf16 before FSDP wrap and disable
669
+ # mixed precision to prevent Accelerate from upcasting bf16→fp32.
670
+ if getattr(self.args, 'fsdp', False):
671
+ model = self.model
672
+ dtype_counts = {}
673
+ for p in model.parameters():
674
+ dt = str(p.dtype)
675
+ dtype_counts[dt] = dtype_counts.get(dt, 0) + 1
676
+
677
+ if len(dtype_counts) > 1:
678
+ logger.info(
679
+ "[HYV4 Patch 4] Mixed dtypes detected before FSDP wrap: %s. "
680
+ "Casting all parameters to bfloat16.",
681
+ dtype_counts
682
+ )
683
+ for p in model.parameters():
684
+ if p.dtype != torch.bfloat16:
685
+ p.data = p.data.to(torch.bfloat16)
686
+ else:
687
+ logger.info(
688
+ "[HYV4 Patch 4] All params already uniform dtype: %s",
689
+ dtype_counts
690
+ )
691
+
692
+ # Disable FSDP mixed precision to prevent Accelerate from
693
+ # upcasting bf16 params to fp32 during wrap. Since all params
694
+ # are already bf16 (cast above), we don't need FSDP's mixed
695
+ # precision policy. This saves ~2x GPU memory per shard.
696
+ try:
697
+ if hasattr(self, 'accelerator'):
698
+ # Level 1: Set accelerator state _mixed_precision to "no"
699
+ # Note: mixed_precision is a property, must set _mixed_precision
700
+ if hasattr(self.accelerator, 'state'):
701
+ old_mp = getattr(self.accelerator.state, '_mixed_precision', None)
702
+ self.accelerator.state._mixed_precision = "no"
703
+ logger.info(
704
+ "[HYV4 Patch 4] Set accelerator.state._mixed_precision='no' "
705
+ "(was: %s) to prevent bf16→fp32 upcast.", old_mp
706
+ )
707
+ # Level 2: Clear the FSDP plugin's mixed_precision_policy
708
+ fsdp_plugin = getattr(self.accelerator.state, 'fsdp_plugin', None)
709
+ if fsdp_plugin is not None:
710
+ if hasattr(fsdp_plugin, 'mixed_precision_policy'):
711
+ fsdp_plugin.mixed_precision_policy = None
712
+ if hasattr(fsdp_plugin, 'kwargs'):
713
+ fsdp_plugin.kwargs.pop('mixed_precision', None)
714
+ logger.info(
715
+ "[HYV4 Patch 4] Cleared fsdp_plugin mixed_precision_policy."
716
+ )
717
+ else:
718
+ logger.warning(
719
+ "[HYV4 Patch 4] self.accelerator not found, cannot disable mixed precision."
720
+ )
721
+ # Level 3: Environment variable (for any lazy initialization)
722
+ os.environ["ACCELERATE_MIXED_PRECISION"] = "no"
723
+ except Exception as e:
724
+ logger.warning(
725
+ "[HYV4 Patch 4] Failed to disable mixed precision: %s", e
726
+ )
727
+
728
+ return _orig_prepare_for_training(self, *args, **kwargs)
729
+
730
+ Trainer._prepare_for_training = _patched_prepare_for_training
731
+ logger.info(
732
+ "HYV4 Patch 4 applied: unify model dtype to bf16 before FSDP wrap."
733
+ )
734
+ except (ImportError, AttributeError) as e:
735
+ logger.warning("[HYV4 Patch 4] Could not apply dtype patch: %s", e)
736
+
737
+ _apply_fsdp_dtype_patch()
738
+
739
+ # Patch 5: Fix CustomSeq2SeqTrainer.create_optimizer signature
740
+ # New transformers (5.15+) calls self.create_optimizer(model) with a model arg,
741
+ # but LLaMA Factory's CustomSeq2SeqTrainer.create_optimizer(self) doesn't accept it.
742
+ def _apply_create_optimizer_patch():
743
+ """Monkey-patch CustomSeq2SeqTrainer.create_optimizer to accept optional model arg."""
744
+ try:
745
+ from llamafactory.train.sft.trainer import CustomSeq2SeqTrainer
746
+
747
+ _orig_create_optimizer = CustomSeq2SeqTrainer.create_optimizer
748
+
749
+ def _patched_create_optimizer(self, model=None):
750
+ return _orig_create_optimizer(self)
751
+
752
+ CustomSeq2SeqTrainer.create_optimizer = _patched_create_optimizer
753
+ logger.info(
754
+ "HYV4 Patch 5 applied: CustomSeq2SeqTrainer.create_optimizer "
755
+ "now accepts optional model argument for transformers >= 5.15."
756
+ )
757
+ except (ImportError, AttributeError) as e:
758
+ logger.warning("[HYV4 Patch 5] Could not apply create_optimizer patch: %s", e)
759
+
760
+ _apply_create_optimizer_patch()
761
+
762
+ # Patch 2 (tokenizer copy) is applied via HYV4PatchCallback during training.
763
+ # Users should add HYV4PatchCallback to their Trainer callbacks.
764
+
765
+ logger.info(
766
+ "HYV4 patches module loaded. Remember to add HYV4PatchCallback to "
767
+ "your Trainer callbacks for tokenizer file copy on save."
768
+ )
finetune/llama_factory_support/hy_v4_template.py ADDED
@@ -0,0 +1,50 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ HYV4 chat template registration for LLaMA Factory.
3
+
4
+ Usage:
5
+ 1. Copy this file's register_template block into LLaMA Factory's
6
+ src/llamafactory/data/template.py (for upstream MR).
7
+ 2. Or import this module before training to register at runtime:
8
+ import hy_v4_template
9
+ """
10
+
11
+ from llamafactory.data.template import ReasoningTemplate, register_template
12
+ from llamafactory.data.formatter import EmptyFormatter, StringFormatter
13
+
14
+
15
+ # ---------------------------------------------------------------------------
16
+ # HYV4 (MoE, pure text) chat template
17
+ #
18
+ # Token format (from chat_template.jinja & tokenizer_config.json):
19
+ # Each turn: <|hy_start:opensource|>{role}<|hy_middle:opensource|>{content}<|hy_end:opensource|>
20
+ # BOS: <|hy_start:opensource|> (token ID 120000)
21
+ # Middle: <|hy_middle:opensource|> (token ID 120001)
22
+ # EOS: <|hy_end:opensource|> (token ID 120025)
23
+ #
24
+ # Loss mask: only compute loss on assistant content (including eos).
25
+ #
26
+ # Reasoning / Slow-thinking support:
27
+ # This template uses ReasoningTemplate with thought_words so that LLaMA
28
+ # Factory can correctly mask think-tag tokens during loss computation.
29
+ # - thought_words: ("<think:opensource>", "</think:opensource>")
30
+ # - enable_thinking: set globally via data_args.enable_thinking (default True)
31
+ #
32
+ # IMPORTANT: To train slow-thinking (chain-of-thought) behaviour, your
33
+ # training data must include <think:opensource>...</think:opensource> tags
34
+ # inside the assistant content. If your data does NOT contain think tags,
35
+ # the model will only learn fast-thinking (direct answer) mode.
36
+ # The `reasoning_effort` field from the API is NOT used by LLaMA Factory;
37
+ # slow-vs-fast is determined solely by the presence of think tags in data.
38
+ # ---------------------------------------------------------------------------
39
+
40
+ register_template(
41
+ name="hy_v4",
42
+ template_class=ReasoningTemplate,
43
+ format_user=StringFormatter(slots=["<|hy_start:opensource|>user<|hy_middle:opensource|>{{content}}<|hy_end:opensource|>"]),
44
+ format_assistant=StringFormatter(slots=["<|hy_start:opensource|>assistant<|hy_middle:opensource|>{{content}}<|hy_end:opensource|>"]),
45
+ format_system=StringFormatter(slots=["<|hy_start:opensource|>system<|hy_middle:opensource|>{{content}}<|hy_end:opensource|>"]),
46
+ format_prefix=EmptyFormatter(slots=[]),
47
+ thought_words=("<think:opensource>", "</think:opensource>"),
48
+ stop_words=["<|hy_end:opensource|>"],
49
+ efficient_eos=False,
50
+ )
finetune/llama_factory_support/train_hy_v4.py ADDED
@@ -0,0 +1,99 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ LLaMA Factory training entry-point wrapper.
3
+
4
+ This script:
5
+ 1. Registers the chat template
6
+ 2. Applies all monkey-patches (checkpoint key rename, dtype fix, etc.)
7
+ 3. Injects HYV4PatchCallback into the training loop
8
+ 4. Calls run_exp() to start LLaMA Factory training
9
+
10
+ """
11
+
12
+ import sys
13
+ import os
14
+
15
+ # Add current directory to path so patches can be imported
16
+ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
17
+
18
+ # Step 1: Register HYV4 template (must be before training starts)
19
+ import hy_v4_template # noqa: F401
20
+
21
+ # Step 2: Apply checkpoint key rename patch (must be before model loading)
22
+ import hy_v4_patches # noqa: F401
23
+
24
+ # Step 3: Inject HYV4PatchCallback into LLaMA Factory's training flow
25
+ from llamafactory.train.sft.workflow import run_sft as _orig_run_sft
26
+
27
+
28
+ def _patched_run_sft(
29
+ model_args, data_args, training_args,
30
+ finetuning_args, generating_args, callbacks=None
31
+ ):
32
+ """Wrap run_sft to inject HYV4PatchCallback."""
33
+ if callbacks is None:
34
+ callbacks = []
35
+
36
+ # Determine tokenizer directory for the save callback
37
+ tokenizer_dir = getattr(model_args, "model_name_or_path", None)
38
+ callbacks.append(
39
+ hy_v4_patches.HYV4PatchCallback(tokenizer_dir=tokenizer_dir)
40
+ )
41
+
42
+ return _orig_run_sft(
43
+ model_args, data_args, training_args,
44
+ finetuning_args, generating_args,
45
+ callbacks=callbacks
46
+ )
47
+
48
+
49
+ # Monkey-patch the SFT workflow
50
+ import llamafactory.train.sft.workflow as _sft_wf
51
+ _sft_wf.run_sft = _patched_run_sft
52
+
53
+
54
+ def _apply_skip_grad_norm_patch():
55
+ """Skip grad norm computation for DeepSpeed ZeRO-3 + CPU offload.
56
+
57
+ Under ZeRO-3 + CPU offload, DeepSpeed's complete_grad_norm_calculation
58
+ all-gathers every gradient on CPU and does an ALLREDUCE to compute the
59
+ global L2 norm. For a 770B model this is extremely slow and can cause
60
+ NCCL timeout/deadlock at optimizer step.
61
+
62
+ When max_grad_norm=0 (no clipping), we fully skip the norm computation
63
+ by patching _get_norm_groups to return 0.0 immediately.
64
+ """
65
+ import torch
66
+ try:
67
+ from deepspeed.runtime.zero.stage3 import DeepSpeedZeroOptimizer_Stage3
68
+
69
+ def _skip_get_norm_groups(self):
70
+ return [torch.tensor(0.0)]
71
+
72
+ DeepSpeedZeroOptimizer_Stage3._get_norm_groups = _skip_get_norm_groups
73
+ print("[HYV4 Patch] Patched DeepSpeedZeroOptimizer_Stage3._get_norm_groups "
74
+ "to skip grad norm computation (max_grad_norm=0).", flush=True)
75
+ except ImportError:
76
+ # DeepSpeed not available, skip
77
+ pass
78
+
79
+
80
+ def main():
81
+ """Entry point: called by torchrun in each worker process.
82
+
83
+ Since train_lf.sh launches us via torchrun directly, all patches
84
+ (template registration, checkpoint key rename, SFT callback injection)
85
+ are already applied in this process. We just call run_exp() to start
86
+ training — no need to go through the CLI launcher.
87
+ """
88
+ # Apply grad norm skip patch for DeepSpeed ZeRO-3 full SFT.
89
+ # This must be done before Trainer creates the DeepSpeed engine.
90
+ # The patch is safe even when not using DeepSpeed (it's a no-op if
91
+ # DeepSpeed is not imported or ZeRO-3 is not used).
92
+ _apply_skip_grad_norm_patch()
93
+
94
+ from llamafactory.train.tuner import run_exp
95
+ run_exp()
96
+
97
+
98
+ if __name__ == "__main__":
99
+ main()
finetune/llama_factory_support/train_lf.sh ADDED
@@ -0,0 +1,105 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/bin/bash
2
+
3
+ set -euo pipefail
4
+
5
+ # -------------------- Network Configuration --------------------
6
+ NET_TYPE="high"
7
+ export NCCL_DEBUG=WARN
8
+ export NCCL_P2P_LEVEL=NVL
9
+ export NCCL_IB_TIMEOUT=24
10
+ export NCCL_NVLS_ENABLE=0
11
+ export NCCL_MPI_PROFILE_PRIMS_ENABLE=0
12
+ export CUDA_DEVICE_MAX_CONNECTIONS=1
13
+ export TORCH_NCCL_HEARTBEAT_TIMEOUT_SEC=3600
14
+ export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,max_split_size_mb:128
15
+ if [[ "${NET_TYPE}" = "low" ]]; then
16
+ export NCCL_SOCKET_IFNAME=eth1
17
+ export NCCL_IB_GID_INDEX=3
18
+ export NCCL_IB_HCA=mlx5_2:1
19
+ export NCCL_IB_SL=3
20
+ export NCCL_CHECK_DISABLE=1
21
+ export NCCL_P2P_DISABLE=0
22
+ export NCCL_LL_THRESHOLD=16384
23
+ export NCCL_IB_CUDA_SUPPORT=1
24
+ else
25
+ export NCCL_IB_GID_INDEX=3
26
+ export NCCL_IB_SL=3
27
+ export NCCL_CHECK_DISABLE=1
28
+ export NCCL_P2P_DISABLE=0
29
+ export NCCL_IB_DISABLE=0
30
+ export NCCL_LL_THRESHOLD=16384
31
+ export NCCL_IB_CUDA_SUPPORT=1
32
+ export NCCL_SOCKET_IFNAME=bond1
33
+ export UCX_NET_DEVICES=bond1
34
+ export NCCL_IB_HCA=mlx5_bond_1,mlx5_bond_5,mlx5_bond_3,mlx5_bond_7,mlx5_bond_4,mlx5_bond_8,mlx5_bond_2,mlx5_bond_6
35
+ export NCCL_COLLNET_ENABLE=0
36
+ export SHARP_COLL_ENABLE_SAT=0
37
+ export NCCL_NET_GDR_LEVEL=2
38
+ export NCCL_IB_QPS_PER_CONNECTION=4
39
+ export NCCL_IB_TC=160
40
+ export NCCL_PXN_DISABLE=1
41
+ fi
42
+
43
+ export DISABLE_VERSION_CHECK=1
44
+
45
+ # -------------------- Node Configuration --------------------
46
+ export HOST_GPU_NUM=8
47
+ # IP list, comma separated. e.g. "10.0.0.1,10.0.0.2" or single node "127.0.0.1"
48
+ export IP_LIST=${IP_LIST:-"127.0.0.1"}
49
+
50
+ MASTER_PORT=${MASTER_PORT:-29500}
51
+
52
+ IFS=',' read -ra IP_ARRAY <<< "$IP_LIST"
53
+ NODES=${#IP_ARRAY[@]}
54
+ MASTER_ADDR=${IP_ARRAY[0]}
55
+
56
+ # -------------------- Paths --------------------
57
+ SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
58
+ YAML_FILE="${YAML_FILE:-hy_v4_full_sft.yaml}"
59
+ # If YAML_FILE is not an absolute path, resolve it relative to SCRIPT_DIR
60
+ if [[ "${YAML_FILE}" != /* ]]; then
61
+ YAML_FILE="${SCRIPT_DIR}/${YAML_FILE}"
62
+ fi
63
+ ENTRY_SCRIPT="${SCRIPT_DIR}/train_hy_v4.py"
64
+
65
+ # -------------------- Distributed Environment --------------------
66
+ export MASTER_ADDR="${MASTER_ADDR}"
67
+ export MASTER_PORT="${MASTER_PORT}"
68
+ export NNODES="${NODES}"
69
+
70
+ # Prevent Accelerate/FSDP from upcasting bf16 parameters back to fp32
71
+ export ACCELERATE_MIXED_PRECISION=no
72
+
73
+ if [ ${NODES} -gt 1 ]; then
74
+ # Determine local node rank by matching local IP against IP_LIST
75
+ LOCAL_IP=$(hostname -i | awk '{print $1}')
76
+ NODE_RANK=0
77
+ for i in "${!IP_ARRAY[@]}"; do
78
+ if [[ "${IP_ARRAY[$i]}" == "${LOCAL_IP}" ]]; then
79
+ NODE_RANK=$i
80
+ break
81
+ fi
82
+ done
83
+ export RANK="${NODE_RANK}"
84
+ else
85
+ export RANK=0
86
+ fi
87
+
88
+ echo "============================================"
89
+ echo " HYV4 LLaMA Factory Training"
90
+ echo " Nodes: ${NNODES}, Rank: ${RANK}"
91
+ echo " Master: ${MASTER_ADDR}:${MASTER_PORT}"
92
+ echo " GPUs per node: ${HOST_GPU_NUM}"
93
+ echo " Total GPUs: $((NODES * HOST_GPU_NUM))"
94
+ echo "============================================"
95
+
96
+ # -------------------- Launch --------------------
97
+ # We launch torchrun directly (instead of FORCE_TORCHRUN) so that each
98
+ # worker process runs train_hy_v4.py with all HYV4 patches applied.
99
+ torchrun \
100
+ --nnodes "${NNODES}" \
101
+ --node_rank "${RANK}" \
102
+ --nproc_per_node "${HOST_GPU_NUM}" \
103
+ --master_addr "${MASTER_ADDR}" \
104
+ --master_port "${MASTER_PORT}" \
105
+ "${ENTRY_SCRIPT}" "${YAML_FILE}"
finetune/ms_swift_support/fsdp_hy_v4.json ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_description": "FSDP configuration",
3
+ "_note": "Uses FSDP with full_shard auto_wrap.",
4
+
5
+ "fsdp": "full_shard auto_wrap",
6
+ "fsdp_config": {
7
+ "version": 1,
8
+ "reshard_after_forward": "full_shard",
9
+ "auto_wrap_policy": "TRANSFORMER_BASED_WRAP",
10
+ "transformer_layer_cls_to_wrap": ["HYV4DecoderLayer"],
11
+ "cpu_ram_efficient_loading": true,
12
+ "state_dict_type": "FULL_STATE_DICT",
13
+ "activation_checkpointing": true,
14
+ "use_orig_params": true,
15
+ "sync_module_states": true,
16
+ "forward_prefetch": false,
17
+ "backward_prefetch": "NO_PREFETCH",
18
+ "limit_all_gathers": true
19
+ }
20
+ }
finetune/ms_swift_support/hy_v4_full_sft.yaml ADDED
@@ -0,0 +1,63 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ # ---- Model Settings ----
3
+ model: /path/to/hy_ckpt
4
+ model_type: hy_v4
5
+ template: hy_v4
6
+ torch_dtype: bfloat16
7
+ attn_impl: sdpa
8
+ trust_remote_code: true
9
+
10
+ # ---- Training Method ----
11
+ train_type: full
12
+
13
+ # ---- Dataset Settings ----
14
+ dataset:
15
+ - ../data/example_data.jsonl
16
+ max_length: 256
17
+ lazy_tokenize: true
18
+ dataset_num_proc: 4
19
+
20
+ # ---- Output Settings ----
21
+ output_dir: saves/hy_v4/full/sft_ds
22
+ save_steps: 500
23
+ save_strategy: steps
24
+ save_total_limit: 3
25
+ save_only_model: false
26
+ logging_steps: 1
27
+ report_to: none
28
+
29
+ # ---- Training Hyperparameters ----
30
+ per_device_train_batch_size: 1
31
+ gradient_accumulation_steps: 1
32
+ learning_rate: 1.0e-5
33
+ num_train_epochs: 3.0
34
+ warmup_steps: 0
35
+ lr_scheduler_type: cosine_with_min_lr
36
+ lr_scheduler_kwargs:
37
+ min_lr_rate: 0.1
38
+ bf16: true
39
+
40
+ # ---- DeepSpeed ZeRO-3 + CPU Offload ----
41
+ # Use custom ds_config for fine-grained control over offload settings
42
+ deepspeed: ../deepspeed_support/ds_zero3_offload.json
43
+
44
+ # ---- Optimization ----
45
+ # adamw_torch is required: adafactor's tuple eps is incompatible with
46
+ # DeepSpeed's CPU Adam kernel (causes TypeError: adam_update())
47
+ optim: adamw_torch
48
+ gradient_checkpointing: true
49
+ # use_reentrant=true is required for DeepSpeed ZeRO-3 to avoid
50
+ # CheckpointError due to parameter partition state changes during recompute
51
+ gradient_checkpointing_kwargs:
52
+ use_reentrant: true
53
+ # max_grad_norm=0 skips grad norm computation (patched in hy_v4_swift_patches.py)
54
+ # to avoid NCCL ALLREDUCE timeout on 770B model
55
+ max_grad_norm: 0.0
56
+
57
+ # ---- Distributed Training ----
58
+ ddp_timeout: 180000000
59
+
60
+ # ---- Misc ----
61
+ seed: 42
62
+ ignore_data_skip: true
63
+ resume_from_checkpoint: null
finetune/ms_swift_support/hy_v4_lora_sft.yaml ADDED
@@ -0,0 +1,65 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ # ---- Model Settings ----
3
+ model: /path/to/hy_ckpt
4
+ model_type: hy_v4
5
+ template: hy_v4
6
+ torch_dtype: bfloat16
7
+ attn_impl: sdpa
8
+ trust_remote_code: true
9
+
10
+ # ---- Training Method: LoRA ----
11
+ train_type: lora
12
+ lora_rank: 64
13
+ lora_alpha: 128
14
+ lora_dropout: 0.05
15
+ # HYV4 uses MLA (Multi-head Latent Attention) with these projection names
16
+ target_modules:
17
+ - q_a_proj
18
+ - q_b_proj
19
+ - kv_a_proj_with_mqa
20
+ - kv_b_proj
21
+ - o_proj
22
+
23
+ # ---- Dataset Settings ----
24
+ dataset:
25
+ - ../data/example_data.jsonl
26
+ max_length: 1024
27
+ lazy_tokenize: true
28
+ dataset_num_proc: 4
29
+
30
+ # ---- Output Settings ----
31
+ output_dir: saves/hy_v4/lora/sft_fsdp
32
+ save_steps: 500
33
+ save_strategy: steps
34
+ save_total_limit: 3
35
+ save_only_model: false
36
+ logging_steps: 1
37
+ report_to: none
38
+
39
+ # ---- Training Hyperparameters ----
40
+ per_device_train_batch_size: 1
41
+ gradient_accumulation_steps: 1
42
+ learning_rate: 2.0e-4
43
+ num_train_epochs: 3.0
44
+ warmup_steps: 0
45
+ lr_scheduler_type: cosine_with_min_lr
46
+ lr_scheduler_kwargs:
47
+ min_lr_rate: 0.1
48
+ bf16: true
49
+
50
+ # ---- FSDP Configuration ----
51
+ # ms-swift uses FSDP (PyTorch native). Pass the JSON config file path to --fsdp.
52
+ # FSDP uses activation_checkpointing (configured in JSON) instead of gradient_checkpointing.
53
+ fsdp: fsdp_hy_v4.json
54
+
55
+ # ---- Optimization ----
56
+ optim: adamw_torch
57
+ max_grad_norm: 1.0
58
+
59
+ # ---- Distributed Training ----
60
+ ddp_timeout: 180000000
61
+
62
+ # ---- Misc ----
63
+ seed: 42
64
+ ignore_data_skip: true
65
+ resume_from_checkpoint: null
finetune/ms_swift_support/hy_v4_swift_patches.py ADDED
@@ -0,0 +1,431 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Patches for ms-swift training.
3
+
4
+ Patches auto-applied on import:
5
+ 1. Model & Template registration: Register the custom model_type
6
+ in ms-swift.
7
+ 2. Grad norm skip patch: Skip grad norm computation for ZeRO-3 + CPU offload.
8
+ 3. Memory-efficient model loading: Shard-by-shard loading that
9
+ processes one safetensors shard for memory optimization.
10
+ 4. Fix logging_dir compatibility between ms-swift and transformers 5.x.
11
+ 5. Align FSDP1 dtype / mixed precision behavior with LLaMA Factory.
12
+ 6. Disable _compute_acc to avoid errors during training.
13
+
14
+ Optional (call manually after LoRA is applied):
15
+ - apply_lora_z3_leaf_patch(model): Mark PEFT LoRA wrapper modules as
16
+ ZeRO-3 leaf modules to fix parameter fetch/release scheduling issues.
17
+
18
+ Usage:
19
+ swift sft --custom_register_path hy_v4_swift_patches.py --model /path/to/ckpt ...
20
+ """
21
+
22
+ import os
23
+ import gc
24
+ import json as _json
25
+ import logging
26
+ from typing import Any, Dict
27
+
28
+ import torch
29
+
30
+ logger = logging.getLogger(__name__)
31
+
32
+ # ============================================================================
33
+ # Patch 1: Model & Template Registration
34
+ #
35
+ # ms-swift natively supports hy_v3, but HYV4 has a different model class
36
+ # (HYV4ForCausalLM) and additional iHC modules. We register HYV4 as a
37
+ # custom model_type that reuses the hy_v3 template.
38
+ # ============================================================================
39
+
40
+ from swift.model import register_model, ModelMeta, ModelGroup, Model
41
+ from swift.template import register_template, TemplateMeta
42
+
43
+ # Register hy_v4 template
44
+ # Token format: <|hy_start:opensource|>{role}<|hy_middle:opensource|>{content}<|hy_end:opensource|>
45
+ register_template(
46
+ TemplateMeta(
47
+ template_type='hy_v4',
48
+ prefix=[],
49
+ system_prefix=['<|hy_start:opensource|>system<|hy_middle:opensource|>{{SYSTEM}}<|hy_end:opensource|>'],
50
+ prompt=['<|hy_start:opensource|>user<|hy_middle:opensource|>{{QUERY}}<|hy_end:opensource|><|hy_start:opensource|>assistant<|hy_middle:opensource|>'],
51
+ chat_sep=['<|hy_end:opensource|>'],
52
+ suffix=['<|hy_end:opensource|>'],
53
+ ),
54
+ exist_ok=True,
55
+ )
56
+
57
+ # Register hy_v4 model
58
+ register_model(
59
+ ModelMeta(
60
+ model_type='hy_v4',
61
+ model_groups=[
62
+ ModelGroup([
63
+ Model('Tencent-Hunyuan/Hy4',
64
+ 'Tencent-Hunyuan/Hy4'),
65
+ ]),
66
+ ],
67
+ template='hy_v4',
68
+ architectures=['HYV4ForCausalLM'],
69
+ is_multimodal=False,
70
+ ),
71
+ exist_ok=True,
72
+ )
73
+
74
+ logger.info(
75
+ "[HYV4 Patch 1] Model type 'hy_v4' and template registered in ms-swift."
76
+ )
77
+
78
+ # ============================================================================
79
+ # Patch 2: Skip grad norm computation for DeepSpeed ZeRO-3
80
+ #
81
+ # Under ZeRO-3 + CPU offload, DeepSpeed's complete_grad_norm_calculation
82
+ # all-gathers every gradient on CPU and does an ALLREDUCE to compute the
83
+ # global L2 norm. For a 770B model this causes NCCL timeout/deadlock.
84
+ # When max_grad_norm=0 (no clipping), we skip the norm computation entirely.
85
+ # ============================================================================
86
+
87
+ def _apply_skip_grad_norm_patch():
88
+ """Patch DeepSpeed ZeRO-3 optimizer to skip grad norm computation."""
89
+ try:
90
+ from deepspeed.runtime.zero.stage3 import DeepSpeedZeroOptimizer_Stage3
91
+
92
+ def _skip_get_norm_groups(self):
93
+ return [torch.tensor(0.0)]
94
+
95
+ DeepSpeedZeroOptimizer_Stage3._get_norm_groups = _skip_get_norm_groups
96
+ logger.info(
97
+ "[HYV4 Patch 2] Patched DeepSpeedZeroOptimizer_Stage3._get_norm_groups "
98
+ "to skip grad norm computation."
99
+ )
100
+ except ImportError:
101
+ logger.info("[HYV4 Patch 2] DeepSpeed not available, skipping grad norm patch.")
102
+
103
+
104
+ # ============================================================================
105
+ # Patch 3: Memory-efficient model loading for ZeRO-3 and FSDP1
106
+ #
107
+ # ZeRO-3 path:
108
+ # Ensure transformers sees the DeepSpeed config early enough to activate its
109
+ # native efficient loading path.
110
+ #
111
+ # FSDP1 path:
112
+ # Mirror the previously successful LLaMA Factory behavior:
113
+ # 1. local_rank 0 loads real weights to CPU only
114
+ # 2. all other local ranks create the model on meta device
115
+ # 3. low_cpu_mem_usage stays enabled
116
+ #
117
+ # This avoids loading the full model onto every GPU during `from_pretrained`,
118
+ # which is exactly the failure mode we observed with ms-swift + FSDP.
119
+ # ============================================================================
120
+
121
+ def _apply_shard_loading_patch():
122
+ """Ensure efficient model loading is used for both ZeRO-3 and FSDP1."""
123
+ import sys
124
+ from transformers import AutoConfig, PreTrainedModel
125
+
126
+ _real_orig_from_pretrained = PreTrainedModel.from_pretrained.__func__
127
+
128
+ def _disable_router_logits_if_needed(model):
129
+ if hasattr(model, 'config') and getattr(model.config, 'output_router_logits', False):
130
+ model.config.output_router_logits = False
131
+ print("[HYV4 Patch 3] Disabled output_router_logits.", flush=True)
132
+ return model
133
+
134
+ def _is_fsdp_requested():
135
+ accel_fsdp = str(os.environ.get("ACCELERATE_USE_FSDP", "")).lower()
136
+ if accel_fsdp in {"1", "true", "yes"}:
137
+ return True
138
+ return "--fsdp" in sys.argv
139
+
140
+ def _build_meta_model_for_fsdp(cls, model_path, kwargs):
141
+ config = kwargs.get("config")
142
+ if config is None:
143
+ config = AutoConfig.from_pretrained(
144
+ model_path,
145
+ trust_remote_code=kwargs.get("trust_remote_code", True),
146
+ )
147
+
148
+ init_kwargs = {}
149
+ torch_dtype = kwargs.get("torch_dtype", None)
150
+ if torch_dtype is not None:
151
+ init_kwargs["torch_dtype"] = torch_dtype
152
+ if "attn_implementation" in kwargs:
153
+ init_kwargs["attn_implementation"] = kwargs["attn_implementation"]
154
+ if "experts_implementation" in kwargs:
155
+ init_kwargs["experts_implementation"] = kwargs["experts_implementation"]
156
+
157
+ with torch.device("meta"):
158
+ model = cls._from_config(config, **init_kwargs)
159
+ return _disable_router_logits_if_needed(model)
160
+
161
+ def _fsdp_safe_load(cls, pretrained_model_name_or_path, *args, **kwargs):
162
+ kwargs = dict(kwargs)
163
+ kwargs.setdefault("low_cpu_mem_usage", True)
164
+
165
+ local_rank = int(os.environ.get("LOCAL_RANK", "0"))
166
+ if local_rank != 0:
167
+ logger.info(
168
+ "[HYV4 Patch 3] FSDP mode: local_rank=%d != 0, creating model on meta device. "
169
+ "Weights will be synchronized from rank 0.",
170
+ local_rank,
171
+ )
172
+ return _build_meta_model_for_fsdp(cls, pretrained_model_name_or_path, kwargs)
173
+
174
+ logger.info(
175
+ "[HYV4 Patch 3] FSDP mode: local_rank=0, loading real weights to CPU "
176
+ "with low_cpu_mem_usage enabled."
177
+ )
178
+ kwargs["device_map"] = {"": "cpu"}
179
+ model = _real_orig_from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs)
180
+ return _disable_router_logits_if_needed(model)
181
+
182
+ def _ensure_zero3_config_and_load(cls, pretrained_model_name_or_path, *args, **kwargs):
183
+ """Ensure HfDeepSpeedConfig is set before calling from_pretrained."""
184
+ model_path = pretrained_model_name_or_path
185
+ print(f"[HYV4 Patch 3] _ensure_zero3_config_and_load called with path: {model_path}", flush=True)
186
+
187
+ if not (isinstance(model_path, str) and os.path.isdir(model_path)):
188
+ return _real_orig_from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs)
189
+
190
+ index_file = os.path.join(model_path, "model.safetensors.index.json")
191
+ single_file = os.path.join(model_path, "model.safetensors")
192
+ if not (os.path.isfile(index_file) or os.path.isfile(single_file)):
193
+ return _real_orig_from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs)
194
+
195
+ if _is_fsdp_requested():
196
+ return _fsdp_safe_load(cls, pretrained_model_name_or_path, *args, **kwargs)
197
+
198
+ from transformers.integrations.deepspeed import is_deepspeed_zero3_enabled
199
+
200
+ if is_deepspeed_zero3_enabled():
201
+ print("[HYV4 Patch 3] ZeRO-3 already enabled, using native from_pretrained.", flush=True)
202
+ model = _real_orig_from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs)
203
+ return _disable_router_logits_if_needed(model)
204
+
205
+ ds_config_path = os.environ.get("DEEPSPEED_CONFIG_FILE", None)
206
+ if ds_config_path is None:
207
+ ds_config_path = os.environ.get("DEEPSPEED_CONFIG", None)
208
+
209
+ if ds_config_path is None:
210
+ for i, arg in enumerate(sys.argv):
211
+ if arg == '--deepspeed' and i + 1 < len(sys.argv):
212
+ ds_config_path = sys.argv[i + 1]
213
+ break
214
+
215
+ if ds_config_path is None or not os.path.isfile(ds_config_path):
216
+ print("[HYV4 Patch 3] No DeepSpeed config found, using default from_pretrained.", flush=True)
217
+ model = _real_orig_from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs)
218
+ return _disable_router_logits_if_needed(model)
219
+
220
+ with open(ds_config_path, "r") as f:
221
+ ds_config = _json.load(f)
222
+
223
+ zero_stage = ds_config.get("zero_optimization", {}).get("stage", 0)
224
+ if zero_stage != 3:
225
+ print(f"[HYV4 Patch 3] Not ZeRO-3 (stage={zero_stage}), using default.", flush=True)
226
+ model = _real_orig_from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs)
227
+ return _disable_router_logits_if_needed(model)
228
+
229
+ print(f"[HYV4 Patch 3] Setting HfDeepSpeedConfig for ZeRO-3 native loading: {ds_config_path}", flush=True)
230
+
231
+ from transformers.integrations.deepspeed import HfDeepSpeedConfig
232
+ _ds_config_obj = HfDeepSpeedConfig(ds_config_path)
233
+
234
+ model = _real_orig_from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs)
235
+
236
+ print("[HYV4 Patch 3] Native ZeRO-3 from_pretrained completed.", flush=True)
237
+ return _disable_router_logits_if_needed(model)
238
+
239
+ @classmethod
240
+ def _patched_from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs):
241
+ return _ensure_zero3_config_and_load(cls, pretrained_model_name_or_path, *args, **kwargs)
242
+
243
+ PreTrainedModel.from_pretrained = _patched_from_pretrained
244
+
245
+ logger.info("[HYV4 Patch 3] Loading patch applied for ZeRO-3 and FSDP1.")
246
+
247
+
248
+ # ============================================================================
249
+ # Optional Patch (NOT auto-applied): LoRA z3_leaf marking
250
+ #
251
+ # PEFT wraps target Linear layers with lora.Linear, adding extra sub-modules
252
+ # (base_layer, lora_A, lora_B). This changes the module tree structure and
253
+ # disrupts ZeRO-3's parameter fetch/release scheduling, causing OOM during
254
+ # backward recomputation. By marking these wrappers as z3_leaf, ZeRO-3 treats
255
+ # them as atomic units, restoring correct scheduling.
256
+ #
257
+ # This is NOT auto-applied because it requires the model to have LoRA already
258
+ # applied. Call apply_lora_z3_leaf_patch(model) manually after LoRA setup.
259
+ # ============================================================================
260
+
261
+ def apply_lora_z3_leaf_patch(model):
262
+ """Mark PEFT LoRA wrapper modules as ZeRO-3 leaf modules.
263
+
264
+ This is an OPTIONAL patch. Call manually AFTER LoRA has been applied
265
+ to the model and BEFORE training starts.
266
+ """
267
+ try:
268
+ from deepspeed.utils import set_z3_leaf_module
269
+ from peft.tuners.lora import Linear as LoraLinear
270
+ except ImportError:
271
+ logger.info("[HYV4 Optional] DeepSpeed or PEFT not available, skipping z3_leaf patch.")
272
+ return
273
+
274
+ z3_leaf_count = 0
275
+ for module in model.modules():
276
+ if isinstance(module, LoraLinear):
277
+ set_z3_leaf_module(module, True)
278
+ z3_leaf_count += 1
279
+
280
+ logger.info("[HYV4 Optional] Marked %d LoraLinear modules with _z3_leaf=True.", z3_leaf_count)
281
+
282
+
283
+ # ============================================================================
284
+ # Patch 4: Fix logging_dir compatibility between ms-swift and
285
+ # transformers 5.x
286
+ #
287
+ # ms-swift's SftArguments._add_version() accesses self.logging_dir,
288
+ # expecting it to be inherited from transformers.TrainingArguments as a
289
+ # dataclass field. However, in transformers 5.x, logging_dir has been
290
+ # deprecated and is no longer included in dataclass fields (init=False or
291
+ # removed from __dataclass_fields__). This causes AttributeError.
292
+ #
293
+ # Fix: Monkey-patch _add_version to ensure logging_dir exists before access.
294
+ # ============================================================================
295
+
296
+ def _apply_logging_dir_patch():
297
+ """Fix SftArguments._add_version for transformers 5.x compatibility."""
298
+ try:
299
+ from swift.arguments.sft_args import SftArguments
300
+
301
+ _orig_add_version = SftArguments._add_version
302
+
303
+ def _patched_add_version(self):
304
+ # Ensure logging_dir attribute exists (transformers 5.x removed it
305
+ # from dataclass fields but ms-swift 4.4.2 still accesses it)
306
+ if not hasattr(self, 'logging_dir'):
307
+ self.logging_dir = None
308
+ # Also ensure run_name exists (may also be affected)
309
+ if not hasattr(self, 'run_name'):
310
+ self.run_name = None
311
+ _orig_add_version(self)
312
+
313
+ SftArguments._add_version = _patched_add_version
314
+ logger.info(
315
+ "[HYV4 Patch 4] Patched SftArguments._add_version for "
316
+ "transformers 5.x logging_dir compatibility."
317
+ )
318
+ except (ImportError, AttributeError) as e:
319
+ logger.info("[HYV4 Patch 4] Could not apply logging_dir patch: %s", e)
320
+
321
+
322
+ # ==========================================================================
323
+ # Patch 5: Mirror LLaMA Factory's FSDP1 dtype + mixed precision guard
324
+ #
325
+ # After LoRA injection, adapter weights can stay in fp32 while the base model
326
+ # is bf16. In addition, Accelerate may re-enable mixed precision policies that
327
+ # upcast bf16 shards back to fp32 during FSDP wrapping. Both behaviors increase
328
+ # memory pressure significantly.
329
+ #
330
+ # We align ms-swift with the previously successful LLaMA Factory setup by:
331
+ # 1. unifying floating-point parameters to bf16 before FSDP wrap
332
+ # 2. disabling Accelerate mixed precision for the FSDP plugin
333
+ # ============================================================================
334
+
335
+ def _apply_fsdp_dtype_patch():
336
+ """Unify parameter dtype to bf16 and disable Accelerate FSDP1 mixed precision."""
337
+ try:
338
+ from transformers import Trainer
339
+
340
+ _orig_prepare_for_training = Trainer._prepare_for_training
341
+
342
+ def _patched_prepare_for_training(self, *args, **kwargs):
343
+ if getattr(self.args, 'fsdp', False):
344
+ dtype_counts = {}
345
+ for p in self.model.parameters():
346
+ dt = str(p.dtype)
347
+ dtype_counts[dt] = dtype_counts.get(dt, 0) + 1
348
+
349
+ if len(dtype_counts) > 1:
350
+ logger.info(
351
+ "[HYV4 Patch 5] Mixed dtypes detected before FSDP wrap: %s. "
352
+ "Casting floating-point parameters to bfloat16.",
353
+ dtype_counts,
354
+ )
355
+ for p in self.model.parameters():
356
+ if p.dtype != torch.bfloat16 and p.dtype.is_floating_point:
357
+ p.data = p.data.to(torch.bfloat16)
358
+ else:
359
+ logger.info("[HYV4 Patch 5] Parameter dtypes already uniform: %s", dtype_counts)
360
+
361
+ try:
362
+ if hasattr(self, 'accelerator') and hasattr(self.accelerator, 'state'):
363
+ old_mp = getattr(self.accelerator.state, '_mixed_precision', None)
364
+ self.accelerator.state._mixed_precision = 'no'
365
+ fsdp_plugin = getattr(self.accelerator.state, 'fsdp_plugin', None)
366
+ if fsdp_plugin is not None:
367
+ if hasattr(fsdp_plugin, 'mixed_precision_policy'):
368
+ fsdp_plugin.mixed_precision_policy = None
369
+ if hasattr(fsdp_plugin, 'kwargs') and isinstance(fsdp_plugin.kwargs, dict):
370
+ fsdp_plugin.kwargs.pop('mixed_precision', None)
371
+ logger.info(
372
+ "[HYV4 Patch 5] Disabled Accelerate FSDP1 mixed precision "
373
+ "(previous state: %s).",
374
+ old_mp,
375
+ )
376
+ os.environ['ACCELERATE_MIXED_PRECISION'] = 'no'
377
+ except Exception as e:
378
+ logger.warning("[HYV4 Patch 5] Failed to disable mixed precision: %s", e)
379
+
380
+ return _orig_prepare_for_training(self, *args, **kwargs)
381
+
382
+ Trainer._prepare_for_training = _patched_prepare_for_training
383
+ logger.info("[HYV4 Patch 5] FSDP1 dtype + mixed precision guard applied.")
384
+ except (ImportError, AttributeError) as e:
385
+ logger.info("[HYV4 Patch 5] Could not apply FSDP1 dtype patch: %s", e)
386
+
387
+
388
+ # ============================================================================
389
+ # Patch 6: Disable _compute_acc during training
390
+ #
391
+ # ms-swift computes training accuracy (argmax on logits) at every step.
392
+ # This requires keeping outputs.logits in memory until acc is computed,
393
+ # adding extra GPU/CPU memory pressure. LLaMA Factory and DeepSpeed native
394
+ # scripts do NOT compute training accuracy. Disabling this aligns ms-swift
395
+ # with the other frameworks and reduces memory usage during forward pass.
396
+ # ============================================================================
397
+
398
+ def _apply_disable_compute_acc_patch():
399
+ """Patch _compute_acc to be a no-op during training."""
400
+ try:
401
+ from swift.trainers.mixin import SwiftMixin
402
+
403
+ def _noop_compute_acc(self, outputs, labels, cu_seqlens=None):
404
+ return
405
+
406
+ SwiftMixin._compute_acc = _noop_compute_acc
407
+ print("[HYV4 Patch 6] Disabled _compute_acc to reduce memory usage.", flush=True)
408
+ except (ImportError, AttributeError) as e:
409
+ print(f"[HYV4 Patch 6] Could not apply _compute_acc patch: {e}", flush=True)
410
+
411
+
412
+ # ============================================================================
413
+ # Auto-apply patches on import
414
+ # ============================================================================
415
+
416
+ # Patch 2: Skip grad norm (always safe to apply; no-op if DeepSpeed not used)
417
+ _apply_skip_grad_norm_patch()
418
+
419
+ # Patch 3: Memory-efficient model loading
420
+ _apply_shard_loading_patch()
421
+
422
+ # Patch 4: Fix logging_dir compatibility
423
+ _apply_logging_dir_patch()
424
+
425
+ # Patch 5: Align FSDP1 dtype / mixed precision behavior with LLaMA Factory
426
+ _apply_fsdp_dtype_patch()
427
+
428
+ # Patch 6: Disable _compute_acc
429
+ _apply_disable_compute_acc_patch()
430
+
431
+ logger.info("[HYV4] All ms-swift patches loaded successfully.")
finetune/ms_swift_support/sft_train.sh ADDED
@@ -0,0 +1,101 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/bin/bash
2
+
3
+ set -euo pipefail
4
+
5
+ # -------------------- Network Configuration --------------------
6
+ NET_TYPE="high"
7
+ export NCCL_DEBUG=WARN
8
+ export NCCL_P2P_LEVEL=NVL
9
+ export NCCL_IB_TIMEOUT=24
10
+ export NCCL_NVLS_ENABLE=0
11
+ export NCCL_MPI_PROFILE_PRIMS_ENABLE=0
12
+ export CUDA_DEVICE_MAX_CONNECTIONS=1
13
+ export TORCH_NCCL_HEARTBEAT_TIMEOUT_SEC=3600
14
+ export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,max_split_size_mb:128
15
+ if [[ "${NET_TYPE}" = "low" ]]; then
16
+ export NCCL_SOCKET_IFNAME=eth1
17
+ export NCCL_IB_GID_INDEX=3
18
+ export NCCL_IB_HCA=mlx5_2:1
19
+ export NCCL_IB_SL=3
20
+ export NCCL_CHECK_DISABLE=1
21
+ export NCCL_P2P_DISABLE=0
22
+ export NCCL_LL_THRESHOLD=16384
23
+ export NCCL_IB_CUDA_SUPPORT=1
24
+ else
25
+ export NCCL_IB_GID_INDEX=3
26
+ export NCCL_IB_SL=3
27
+ export NCCL_CHECK_DISABLE=1
28
+ export NCCL_P2P_DISABLE=0
29
+ export NCCL_IB_DISABLE=0
30
+ export NCCL_LL_THRESHOLD=16384
31
+ export NCCL_IB_CUDA_SUPPORT=1
32
+ export NCCL_SOCKET_IFNAME=bond1
33
+ export UCX_NET_DEVICES=bond1
34
+ export NCCL_IB_HCA=mlx5_bond_1,mlx5_bond_5,mlx5_bond_3,mlx5_bond_7,mlx5_bond_4,mlx5_bond_8,mlx5_bond_2,mlx5_bond_6
35
+ export NCCL_COLLNET_ENABLE=0
36
+ export SHARP_COLL_ENABLE_SAT=0
37
+ export NCCL_NET_GDR_LEVEL=2
38
+ export NCCL_IB_QPS_PER_CONNECTION=4
39
+ export NCCL_IB_TC=160
40
+ export NCCL_PXN_DISABLE=1
41
+ fi
42
+
43
+ # -------------------- Node Configuration --------------------
44
+ export HOST_GPU_NUM=8
45
+ # IP list, comma separated. e.g. "10.0.0.1,10.0.0.2" or single node "127.0.0.1"
46
+ export IP_LIST=${IP_LIST:-"127.0.0.1"}
47
+
48
+ MASTER_PORT=${MASTER_PORT:-29500}
49
+
50
+ IFS=',' read -ra IP_ARRAY <<< "$IP_LIST"
51
+ NODES=${#IP_ARRAY[@]}
52
+ MASTER_ADDR=${IP_ARRAY[0]}
53
+
54
+ # -------------------- Paths --------------------
55
+ SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
56
+ YAML_FILE="${YAML_FILE:-hy_v4_full_sft.yaml}"
57
+ PATCHES_FILE="${SCRIPT_DIR}/hy_v4_swift_patches.py"
58
+ CONFIG_FILE="${SCRIPT_DIR}/${YAML_FILE}"
59
+
60
+ # -------------------- Distributed Environment --------------------
61
+ export MASTER_ADDR="${MASTER_ADDR}"
62
+ export MASTER_PORT="${MASTER_PORT}"
63
+ export NNODES="${NODES}"
64
+
65
+ if [ ${NODES} -gt 1 ]; then
66
+ # Determine local node rank by matching local IP against IP_LIST
67
+ LOCAL_IP=$(hostname -i | awk '{print $1}')
68
+ NODE_RANK=0
69
+ for i in "${!IP_ARRAY[@]}"; do
70
+ if [[ "${IP_ARRAY[$i]}" == "${LOCAL_IP}" ]]; then
71
+ NODE_RANK=$i
72
+ break
73
+ fi
74
+ done
75
+ export RANK="${NODE_RANK}"
76
+ else
77
+ export RANK=0
78
+ fi
79
+
80
+ echo "============================================"
81
+ echo " HYV4 ms-swift SFT Training (Full + DeepSpeed)"
82
+ echo " Nodes: ${NNODES}, Rank: ${RANK}"
83
+ echo " Master: ${MASTER_ADDR}:${MASTER_PORT}"
84
+ echo " GPUs per node: ${HOST_GPU_NUM}"
85
+ echo " Total GPUs: $((NODES * HOST_GPU_NUM))"
86
+ echo " Config: ${CONFIG_FILE}"
87
+ echo "============================================"
88
+
89
+ # -------------------- Launch --------------------
90
+ # Add script directory to PYTHONPATH so patches can find local modules
91
+ export PYTHONPATH="${PYTHONPATH:+${PYTHONPATH}:}${SCRIPT_DIR}"
92
+
93
+ torchrun \
94
+ --nnodes "${NNODES}" \
95
+ --node_rank "${RANK}" \
96
+ --nproc_per_node "${HOST_GPU_NUM}" \
97
+ --master_addr "${MASTER_ADDR}" \
98
+ --master_port "${MASTER_PORT}" \
99
+ -m swift.cli.sft \
100
+ --custom_register_path "${PATCHES_FILE}" \
101
+ --config "${CONFIG_FILE}"
finetune/ms_swift_support/sft_train_lora.sh ADDED
@@ -0,0 +1,104 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/bin/bash
2
+
3
+ set -euo pipefail
4
+
5
+ # -------------------- Network Configuration --------------------
6
+ NET_TYPE="high"
7
+ export NCCL_DEBUG=WARN
8
+ export NCCL_P2P_LEVEL=NVL
9
+ export NCCL_IB_TIMEOUT=24
10
+ export NCCL_NVLS_ENABLE=0
11
+ export NCCL_MPI_PROFILE_PRIMS_ENABLE=0
12
+ export CUDA_DEVICE_MAX_CONNECTIONS=1
13
+ export TORCH_NCCL_HEARTBEAT_TIMEOUT_SEC=3600
14
+ export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,max_split_size_mb:128
15
+ if [[ "${NET_TYPE}" = "low" ]]; then
16
+ export NCCL_SOCKET_IFNAME=eth1
17
+ export NCCL_IB_GID_INDEX=3
18
+ export NCCL_IB_HCA=mlx5_2:1
19
+ export NCCL_IB_SL=3
20
+ export NCCL_CHECK_DISABLE=1
21
+ export NCCL_P2P_DISABLE=0
22
+ export NCCL_LL_THRESHOLD=16384
23
+ export NCCL_IB_CUDA_SUPPORT=1
24
+ else
25
+ export NCCL_IB_GID_INDEX=3
26
+ export NCCL_IB_SL=3
27
+ export NCCL_CHECK_DISABLE=1
28
+ export NCCL_P2P_DISABLE=0
29
+ export NCCL_IB_DISABLE=0
30
+ export NCCL_LL_THRESHOLD=16384
31
+ export NCCL_IB_CUDA_SUPPORT=1
32
+ export NCCL_SOCKET_IFNAME=bond1
33
+ export UCX_NET_DEVICES=bond1
34
+ export NCCL_IB_HCA=mlx5_bond_1,mlx5_bond_5,mlx5_bond_3,mlx5_bond_7,mlx5_bond_4,mlx5_bond_8,mlx5_bond_2,mlx5_bond_6
35
+ export NCCL_COLLNET_ENABLE=0
36
+ export SHARP_COLL_ENABLE_SAT=0
37
+ export NCCL_NET_GDR_LEVEL=2
38
+ export NCCL_IB_QPS_PER_CONNECTION=4
39
+ export NCCL_IB_TC=160
40
+ export NCCL_PXN_DISABLE=1
41
+ fi
42
+
43
+ # -------------------- Node Configuration --------------------
44
+ export HOST_GPU_NUM=8
45
+ # IP list, comma separated. e.g. "10.0.0.1,10.0.0.2" or single node "127.0.0.1"
46
+ export IP_LIST=${IP_LIST:-"127.0.0.1"}
47
+
48
+ MASTER_PORT=${MASTER_PORT:-29500}
49
+
50
+ IFS=',' read -ra IP_ARRAY <<< "$IP_LIST"
51
+ NODES=${#IP_ARRAY[@]}
52
+ MASTER_ADDR=${IP_ARRAY[0]}
53
+
54
+ # -------------------- Paths --------------------
55
+ SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
56
+ YAML_FILE="${YAML_FILE:-hy_v4_lora_sft.yaml}"
57
+ PATCHES_FILE="${SCRIPT_DIR}/hy_v4_swift_patches.py"
58
+ CONFIG_FILE="${SCRIPT_DIR}/${YAML_FILE}"
59
+
60
+ # -------------------- Distributed Environment --------------------
61
+ export MASTER_ADDR="${MASTER_ADDR}"
62
+ export MASTER_PORT="${MASTER_PORT}"
63
+ export NNODES="${NODES}"
64
+
65
+ # Prevent Accelerate/FSDP from upcasting bf16 parameters back to fp32
66
+ export ACCELERATE_MIXED_PRECISION=no
67
+
68
+ if [ ${NODES} -gt 1 ]; then
69
+ # Determine local node rank by matching local IP against IP_LIST
70
+ LOCAL_IP=$(hostname -i | awk '{print $1}')
71
+ NODE_RANK=0
72
+ for i in "${!IP_ARRAY[@]}"; do
73
+ if [[ "${IP_ARRAY[$i]}" == "${LOCAL_IP}" ]]; then
74
+ NODE_RANK=$i
75
+ break
76
+ fi
77
+ done
78
+ export RANK="${NODE_RANK}"
79
+ else
80
+ export RANK=0
81
+ fi
82
+
83
+ echo "============================================"
84
+ echo " HYV4 ms-swift SFT Training (LoRA + FSDP)"
85
+ echo " Nodes: ${NNODES}, Rank: ${RANK}"
86
+ echo " Master: ${MASTER_ADDR}:${MASTER_PORT}"
87
+ echo " GPUs per node: ${HOST_GPU_NUM}"
88
+ echo " Total GPUs: $((NODES * HOST_GPU_NUM))"
89
+ echo " Config: ${CONFIG_FILE}"
90
+ echo "============================================"
91
+
92
+ # -------------------- Launch --------------------
93
+ # Add script directory to PYTHONPATH so patches can find local modules
94
+ export PYTHONPATH="${PYTHONPATH:+${PYTHONPATH}:}${SCRIPT_DIR}"
95
+
96
+ torchrun \
97
+ --nnodes "${NNODES}" \
98
+ --node_rank "${RANK}" \
99
+ --nproc_per_node "${HOST_GPU_NUM}" \
100
+ --master_addr "${MASTER_ADDR}" \
101
+ --master_port "${MASTER_PORT}" \
102
+ -m swift.cli.sft \
103
+ --custom_register_path "${PATCHES_FILE}" \
104
+ --config "${CONFIG_FILE}"
finetune/requirements.txt ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ transformers>=5.16.2
2
+ torch>=2.10.0
3
+ torchvision
4
+ torchaudio
5
+ accelerate>=1.11.0
6
+ peft>=0.18.1
7
+ deepspeed>=0.18.7
8
+ sentencepiece
9
+ protobuf
10
+ ninja
11
+ flash-attn
12
+ tensorboard
generation_config.json ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "bos_token_id": 120000,
3
+ "do_sample": true,
4
+ "eos_token_id": 120025,
5
+ "pad_token_id": 120002,
6
+ "temperature": 0.9,
7
+ "top_k": -1,
8
+ "top_p": 1,
9
+ "transformers_version": "5.16.2"
10
+ }
model-00001-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a1466bb060b882d4c87cfc08a9b75e6693eca42dccea62ceb654ab6e1aac146c
3
+ size 19327353080
model-00002-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:68508b2930a22ea7647b82becdddb7424d6b9569446b985f9977a76276faf661
3
+ size 19327353080
model-00003-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:343973567fefac38c5bb05bc731c85c7dcf203a11e6d53b962888cf940bc0518
3
+ size 19327353080
model-00004-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b230b1208a74cd5b287fca5e5864fca113c55a5bafc2d139d2133f479fce3391
3
+ size 19327353080
model-00005-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d672431ffd2b4e8c89234129155fe6835ccee8861bb7c4ffc8b5bd1b08c01f4d
3
+ size 19327353080
model-00006-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4c32701f46b426fcbdd2ce30affedf397349242d292c6aac444fab00695f2b57
3
+ size 19327353080
model-00007-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:316642aaf64a434cc926e4c8225679aca8d173f3007ef00d8f6fa304c138a450
3
+ size 19327353080
model-00008-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c0566210325918eeab2efeb5208802d124260f901725060c42e9c70a4c989d10
3
+ size 19327353080
model-00009-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c1289a27921de91ea47be3c781a69b45b2861fad0aff4beeeb82a169096f1fdd
3
+ size 19327353080
model-00010-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f2c7046ffbfa7ffadf42dc85003f359f792e5b81751a906946e6f3ec53d0ce3c
3
+ size 19327353080
model-00011-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:37577ce86c912e109a720faa23e5a4a09bbad2f527936b14c933fdcf43847ee7
3
+ size 19327353080
model-00012-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e239142bd5cb88a870d6f58b01914210f1e5d95dff40f1d95aa981e56acabdab
3
+ size 19327353080
model-00013-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:df14d26d081bcfe5c892313520a8f21b4d7e337e284c64695915ed4dd3e935b7
3
+ size 19327353080
model-00014-of-00131.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:27e7bddd536feab082d99627b0344f641391a2b1d44abc1c271ceb5cf1f6ca1d
3
+ size 19327353080