初始化项目,由ModelHub XC社区提供模型
Model: kakaocorp/kanana-2-1.3b-instruct Source: Original Platform
This commit is contained in:
36
.gitattributes
vendored
Normal file
36
.gitattributes
vendored
Normal file
@@ -0,0 +1,36 @@
|
||||
*.7z filter=lfs diff=lfs merge=lfs -text
|
||||
*.arrow filter=lfs diff=lfs merge=lfs -text
|
||||
*.bin filter=lfs diff=lfs merge=lfs -text
|
||||
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
||||
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
||||
*.ftz filter=lfs diff=lfs merge=lfs -text
|
||||
*.gz filter=lfs diff=lfs merge=lfs -text
|
||||
*.h5 filter=lfs diff=lfs merge=lfs -text
|
||||
*.joblib filter=lfs diff=lfs merge=lfs -text
|
||||
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
||||
*.mlmodel filter=lfs diff=lfs merge=lfs -text
|
||||
*.model filter=lfs diff=lfs merge=lfs -text
|
||||
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
||||
*.npy filter=lfs diff=lfs merge=lfs -text
|
||||
*.npz filter=lfs diff=lfs merge=lfs -text
|
||||
*.onnx filter=lfs diff=lfs merge=lfs -text
|
||||
*.ot filter=lfs diff=lfs merge=lfs -text
|
||||
*.parquet filter=lfs diff=lfs merge=lfs -text
|
||||
*.pb filter=lfs diff=lfs merge=lfs -text
|
||||
*.pickle filter=lfs diff=lfs merge=lfs -text
|
||||
*.pkl filter=lfs diff=lfs merge=lfs -text
|
||||
*.pt filter=lfs diff=lfs merge=lfs -text
|
||||
*.pth filter=lfs diff=lfs merge=lfs -text
|
||||
*.rar filter=lfs diff=lfs merge=lfs -text
|
||||
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
||||
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
||||
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
||||
*.tar filter=lfs diff=lfs merge=lfs -text
|
||||
*.tflite filter=lfs diff=lfs merge=lfs -text
|
||||
*.tgz filter=lfs diff=lfs merge=lfs -text
|
||||
*.wasm filter=lfs diff=lfs merge=lfs -text
|
||||
*.xz filter=lfs diff=lfs merge=lfs -text
|
||||
*.zip filter=lfs diff=lfs merge=lfs -text
|
||||
*.zst filter=lfs diff=lfs merge=lfs -text
|
||||
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
||||
assets/logo/kanana.png filter=lfs diff=lfs merge=lfs -text
|
||||
186
LICENSE
Normal file
186
LICENSE
Normal file
@@ -0,0 +1,186 @@
|
||||
# KANANA OPEN LICENSE AGREEMENT
|
||||
|
||||
|
||||
|
||||
Kanana Release Date: July 17, 2025
|
||||
|
||||
This KANANA OPEN LICENSE AGREEMENT (this “Agreement”) is made by and between you and Kakao Corp. (“KAKAO”) that governs your use of Kanana Materials that KAKAO provides to you.
|
||||
|
||||
By using, copying, modifying, distributing, performing, or displaying all or part of Kanana Materials, or otherwise accepting the terms and conditions of this Agreement, you agree to be bound by this Agreement. You hereby represent and warrant that (i) you are legally authorized to enter into this Agreement, and (ii) if you are entering into this Agreement on behalf of a legal entity, you have the authority to legally and validly bind such entity.
|
||||
|
||||
---
|
||||
|
||||
## 1. Definition
|
||||
|
||||
|
||||
|
||||
* **“Agreement”** means the terms and conditions for use, copying, distribution and modification of Kanana Materials as set forth herein.
|
||||
|
||||
|
||||
* **“KAKAO”** means Kakao Corp.
|
||||
|
||||
|
||||
* **“You”** means an individual or legal entity that enters into this Agreement with KAKAO and exercises its rights hereunder or uses Kanana Materials for any purpose. If you enter into this Agreement on behalf of a legal entity, “you” shall include such entity.
|
||||
|
||||
|
||||
* **“Kanana”** means the basic large-scale language model, software, and algorithms distributed by KAKAO under this Agreement, including parameters (such as Model Weights and optimizer status), machine learning model codes, inference/learning/fine-tuning codes, and other related elements.
|
||||
|
||||
|
||||
* **“Documentation”** means the specifications, manuals, and other documentation accompanying Kanana distributed by KAKAO.
|
||||
|
||||
|
||||
* **“Kanana Materials”** means, collectively, Kanana and Documentation, including any portions or components thereof.
|
||||
|
||||
|
||||
* **“Outputs”** means information content generated by operating or otherwise using Kanana Materials.
|
||||
|
||||
|
||||
* **“Derivative Works”** means (i) any modifications to Kanana, (ii) any work of authorship based on Kanana, or (iii) any other designed machine learning models that either directly use the patterns of Model Weights, parameters, operations, and/or outputs or incorporate a substantial part of Kanana’s performance or functional characteristics through methods including, but not limited to, transfer learning, fine-tuning, or knowledge distillation. This includes distillation methods using Kanana’s intermediate data representations or a method based on the synthetic data outputs generated by Kanana; *provided, however*, that Outputs shall not be deemed to be Derivative Works.
|
||||
|
||||
|
||||
* **“Model Weights”** means a set of numerical parameter values generated during Kanana’s learning process, representing the result of substantial investment and effort by KAKAO.
|
||||
|
||||
|
||||
|
||||
---
|
||||
|
||||
## 2. Grant of License and Use Policy
|
||||
|
||||
|
||||
|
||||
**2.1 Grant of License.** Subject to the terms and conditions of this Agreement, you are granted a non-exclusive, worldwide, non-transferrable, royalty-free limited license under KAKAO’s intellectual property or other rights owned by KAKAO that enables you to access, download, install, copy, use, reproduce, distribute, create Derivative Works of, and make modifications to Kanana Materials.
|
||||
|
||||
**2.2 Policy on Prohibited Use.** Your use of Kanana Materials and Derivative Works must comply with applicable laws and regulations and adhere to KAKAO’s Guidelines For Responsible AI, which is hereby incorporated into this Agreement.
|
||||
|
||||
**2.3** This Agreement applies solely to Kanana-*** and shall not apply to any other models distributed by KAKAO under separate licenses. Licenses applicable to such other models shall not apply to Kanana-***.
|
||||
|
||||
**2.4** The license terms applicable to a specific version of Kanana applies exclusively to that version and shall not extend to any other versions. Each version shall be deemed as an independent and separate work of authorship.
|
||||
|
||||
**2.5** You may use each version of Kanana only in accordance with the license terms expressly specified for that version, and you shall not claim that the license terms applicable to one version apply to any other version.
|
||||
|
||||
**2.6** You shall not combine different versions of Kanana versions that are subject to different license terms in order to circumvent any applicable license terms.
|
||||
|
||||
---
|
||||
|
||||
## 3. Redistribution
|
||||
|
||||
|
||||
|
||||
**3.1** You may copy, distribute or disclose Kanana, Derivative Works, or any products or services that contain Kanana or Derivative Works; provided, however, that you shall:
|
||||
|
||||
* **(i)** incorporate the compliance obligation set forth in the Policy on Prohibited Use provision of Section 2.2 in any agreement for use and distribution and notify subsequent users that such use restrictions apply;
|
||||
|
||||
|
||||
* **(ii)** provide any recipients of Kanana Materials or Derivative Works a copy of this Agreement;
|
||||
|
||||
|
||||
* **(iii)** expressly indicate in any files you have modified that it has been modified by you;
|
||||
|
||||
|
||||
* **(iv)** include a “Notice” text file that includes the following notice:
|
||||
> “Kanana is licensed in accordance with the Kanana Open License Agreement. Copyright © KAKAO Corp. All Rights Reserved.”; and
|
||||
>
|
||||
>
|
||||
|
||||
|
||||
* **(v)** clearly display the phrase **“Powered by Kanana”** on related websites, user interfaces, blog posts, introduction pages, or product documentation in a manner that is easily recognizable to users. In addition, if you use Kanana Materials or their outputs to create, train, improve, or enhance other AI models and distribute them, you must include **‘Kanana’** as a prefix to the name of such AI models.
|
||||
|
||||
|
||||
|
||||
**3.2** You may add your own copyright statement to your modifications of Kanana Materials and may provide additional or different license terms and conditions; provided, however, that such additional or different license terms and conditions shall not violate or conflict with any provisions of this Agreement.
|
||||
|
||||
---
|
||||
|
||||
## 4. Additional Commercial Terms
|
||||
|
||||
|
||||
|
||||
**4.1** If you wish to engage in any of the following activities using Kanana Materials or any Derivative Works, you must obtain a separate commercial license expressly granted by KAKAO:
|
||||
|
||||
* **(i)** Offering or (re)selling to third parties access to Kanana Materials or any Derivative Works through API, cloud platforms, or other remote access services;
|
||||
|
||||
|
||||
* **(ii)** Offering or (re)selling to third parties Kanana Materials or any Derivative Works in whole or in part, as part of a system integration (SI) or on-premise deployment solution; or
|
||||
|
||||
|
||||
* **(iii)** Offering or (re)selling to third parties Kanana Materials or any Derivative Works embedded in on-device domains.
|
||||
|
||||
|
||||
|
||||
**4.2** For clarity, unless your activities or conditions fall within those specified in Section 4.1 above, you may use Kanana Materials or any Derivative Works for the development and operation of your own services without obtaining a commercial license from KAKAO.
|
||||
|
||||
**4.3** The grant of any commercial license under Section 4.1 shall be at KAKAO’s sole discretion.
|
||||
|
||||
---
|
||||
|
||||
## 5. Outputs
|
||||
|
||||
|
||||
|
||||
KAKAO will not claim any rights to Outputs you generate using Kanana Materials. You shall be solely responsible for Outputs and the use thereof.
|
||||
|
||||
---
|
||||
|
||||
## 6. Disclaimer of Warranty
|
||||
|
||||
|
||||
|
||||
UNLESS REQUIRED BY LAW, KANANA MATERIALS ARE PROVIDED ON AN “AS IS” BASIS, AND KAKAO DISCLAIMS ALL WARRANTIES OF ANY KIND, BOTH EXPRESS AND IMPLIED, INCLUDING, WITHOUT LIMITATION, ANY WARRANTIES OF TITLE, NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
|
||||
---
|
||||
|
||||
## 7. Limitation on Liability
|
||||
|
||||
|
||||
|
||||
UNLESS REQUIRED BY LAW, IN NO EVENT SHALL KAKAO BE LIABLE TO YOU FOR DAMAGES, INCLUDING ANY DIRECT, INDIRECT, SPECIAL, CONSEQUENTIAL, INCIDENTAL, AND PUNITIVE DAMAGES OF ANY CHARACTER ARISING OUT OF THE USE OR INABILITY TO USE KANANA MATERIALS, DERIVATIVE WORKS, OR OUTPUTS, EVEN IF KAKAO HAS BEEN ADVISED OF THE POSSIBILITY OF SUCH DAMAGES.
|
||||
|
||||
---
|
||||
|
||||
## 8. Indemnification
|
||||
|
||||
|
||||
|
||||
You shall indemnify and hold KAKAO harmless from and against any and all claims that may be filed by a third party as a result of your infringement of any third party’s rights or violation of any applicable law, to the extent caused by your use or distribution of Kanana Materials, Derivative Works, or Outputs; provided, however, that the foregoing shall not apply to claims resulting from KAKAO’s willful or gross negligence.
|
||||
|
||||
---
|
||||
|
||||
## 9. Intellectual Property
|
||||
|
||||
|
||||
|
||||
**9.1** This Agreement does not grant you any rights to use KAKAO’s trademarks, service marks, or product names. However, on a limited basis and solely for the purpose of complying with Section 3.1(v), KAKAO authorizes you to use the Kanana trademark, provided that KAKAO may require you to discontinue such use at any time if you impair the value of the Kanana trademark. For clarity, such limited license for trademark use shall not be interpreted as KAKAO’s endorsement, sponsorship, or guarantee of, or the assumption of any liability for, your Derivative Works, Outputs, or any associated services.
|
||||
|
||||
**9.2** KAKAO retains ownership of Kanana Materials and Derivative Works created by KAKAO, but you will retain ownership of any Derivative Works and modifications made by you.
|
||||
|
||||
**9.3** If you bring any legal action or proceeding against KAKAO or a third party alleging that the Kanana Materials, Derivative Works, or Outputs infringe your intellectual property rights, your rights under this Agreement shall automatically terminate as of the date such action is filed.
|
||||
|
||||
**9.4** You acknowledge that Model Weights are a valuable asset of KAKAO. You shall not extract, copy, distribute, modify Model Weights or use them to train new models, except as expressly permitted under this Agreement.
|
||||
|
||||
**9.5** The protections under this Agreement apply to all components of Kanana Materials (irrespective of whether it is recognized as a work of authorship), including, but not limited to, Model Weights, parameters, algorithms, or structures. You may exercise your rights in these components only to the extent expressly permitted under this Agreement.
|
||||
|
||||
---
|
||||
|
||||
## 10. Term and Termination
|
||||
|
||||
|
||||
|
||||
The term of this Agreement will commence upon your acceptance of this Agreement or access to Kanana Materials and will continue in full force and effect until terminated in accordance with the terms and conditions herein. KAKAO may terminate this Agreement if you are in breach of any term or condition of this Agreement. Upon termination of this Agreement, you shall delete and cease use of Kanana Materials and Derivative Works. Sections 5, 6, 7, 8, 10 and 11 shall survive the termination of this Agreement.
|
||||
|
||||
---
|
||||
|
||||
## 11. Governing Law and Arbitration
|
||||
|
||||
|
||||
|
||||
**11.1** This Agreement will be governed and construed under the laws of the Republic of Korea, without regard to its conflicts of laws principles.
|
||||
|
||||
**11.2** Any disputes arising out of or in connection with this Agreement shall be finally settled by arbitration in accordance with the International Arbitration Rules of the Korean Commercial Arbitration Board. The number of arbitrators shall be one. The seat, or legal place, of arbitral proceedings shall be Seoul, Republic of Korea. The language to be used in the arbitral proceedings shall be English. Either party may seek interim or provisional relief from a court of competent jurisdiction, which shall not be considered a waiver of any provision in this Section. The arbitral tribunal also has the authority to issue orders for interim or provisional relief.
|
||||
|
||||
---
|
||||
|
||||
## 12. No Waiver
|
||||
|
||||
|
||||
|
||||
KAKAO’s failure or delay in exercising any of its rights under this Agreement shall not constitute a waiver of such rights.
|
||||
246
README.md
Normal file
246
README.md
Normal file
@@ -0,0 +1,246 @@
|
||||
---
|
||||
|
||||
library_name: transformers
|
||||
license: other
|
||||
license_name: "kanana-open-license"
|
||||
license_link: https://huggingface.co/kakaocorp/kanana-2-1.3b-instruct/blob/main/LICENSE
|
||||
pipeline_tag: text-generation
|
||||
model_id: kakaocorp/kanana-2-1.3b-instruct
|
||||
repo: kakaocorp/kanana-2-1.3b-instruct
|
||||
developers: Kanana LLM
|
||||
base_model:
|
||||
|
||||
- kakaocorp/kanana-2-1.3b-base
|
||||
|
||||
---
|
||||
|
||||
|
||||
<p align="center">
|
||||
<img src="./assets/logo/kanana.png" width="60%" alt="Kanana">
|
||||
</p>
|
||||
|
||||
<p align="center">
|
||||
🤗 <a href="https://huggingface.co/collections/kakaocorp/kanana-2-slm">HF Models</a> | 📕 <a href="https://tech.kakao.com/posts/826">Blog</a>
|
||||
</p>
|
||||
<br><br>
|
||||
|
||||
## News 🔥
|
||||
|
||||
- `2026/07/27`: 🤗 Released `kanana-2-3b`, `kanana-2-1.3b` HF model weights.
|
||||
- `2026/07/27`: 📕 Published a blog post about the development of the `Kanana-2 SLM` series.
|
||||
|
||||
# Introduction
|
||||
|
||||
We present Kanana-2 SLM, **Kakao's second series of Small Language Models (SLMs)**, designed to deliver strong language capabilities while remaining compact and efficient for practical deployment. The series includes 3B, 1.3B, and 0.9B models. This release publicly includes the 3B model and the compressed 1.3B model, providing a balance between capability and efficiency for a wide range of applications.
|
||||
Kanana-2-3B was **pretrained from scratch on TPU clusters** and further improved through post-training with **supervised fine-tuning and reinforcement learning**, resulting in strong instruction-following and reasoning capabilities.
|
||||
Kanana-2-1.3B models are derived from Kanana-2-3B through a **cascade pruning and distillation pipeline**. To further improve deployment efficiency, they adopt **Sliding Window Attention (SWA)**, enabling memory-efficient long-context inference with support for context lengths of up to **32K tokens** while substantially reducing KV-cache memory requirements.
|
||||
|
||||
The Kanana-2 SLM release consists of the following four publicly available models:
|
||||
|
||||
- **Kanana-2-3B-Base** — 3B pretrained base
|
||||
- **Kanana-2-3B-Instruct** — instruction-tuned 3B model
|
||||
- **Kanana-2-1.3B-Base** — compressed 1.3B pretrained base
|
||||
- **Kanana-2-1.3B-Instruct** — instruction-tuned 1.3B model for on-device deployment
|
||||
|
||||
> [!NOTE]
|
||||
> No Kakao user data was used for either pre-training or post-training.
|
||||
|
||||
|
||||
|
||||
# Highlights
|
||||
|
||||
- **Cascade Pruning & Distillation**: Kanana-2-1.3B is built by progressively compressing Kanana-2-3B-Base (3B → 2B → 1.3B → 0.9B) through a cascade pruning and distillation pipeline.
|
||||
- **Sliding Window Attention (SWA)**: Uses a 3:1 hybrid layout of sliding-window and full-attention layers. A sliding-window size of 1024 reduces per-token KV-cache reads, cutting KV-cache usage by up to ~72.7% at a 32K context length compared to a full-attention-only model. YaRN is applied to full-attention layers, while SWA layers retain RoPE, preserving long-range context without sacrificing local-attention efficiency.
|
||||
- **Kanana-2 tokenizer**: Improves Korean tokenization efficiency by over 30% compared to the previous generation.
|
||||
- **Long context**: Natively supports context lengths of up to 32,768 tokens.
|
||||
|
||||
|
||||
|
||||
## Model Downloads
|
||||
|
||||
|
||||
| **Model** | **Download** |
|
||||
| ---------------------- | ------------------------------------------------------------------------- |
|
||||
| kanana-2-3b-base | [🤗 HuggingFace](https://huggingface.co/kakaocorp/kanana-2-3b-base) |
|
||||
| kanana-2-3b-instruct | [🤗 HuggingFace](https://huggingface.co/kakaocorp/kanana-2-3b-instruct) |
|
||||
| kanana-2-1.3b-base | [🤗 HuggingFace](https://huggingface.co/kakaocorp/kanana-2-1.3b-base) |
|
||||
| kanana-2-1.3b-instruct | [🤗 HuggingFace](https://huggingface.co/kakaocorp/kanana-2-1.3b-instruct) |
|
||||
|
||||
|
||||
|
||||
|
||||
## Performance
|
||||
|
||||
|
||||
|
||||
### Base model evaluation results
|
||||
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>Benchmark</th>
|
||||
<th>Metric</th>
|
||||
<th>Shot</th>
|
||||
<th>kanana-2-3b-base</th>
|
||||
<th>kanana-2-1.3b-base</th>
|
||||
<th>Qwen3-1.7B-Base</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr><td align="center" colspan="6"><b>General Tasks</b></td></tr>
|
||||
<tr><td>MMLU</td><td>acc</td><td>5</td><td>62.77</td><td>56.49</td><td>62.32</td></tr>
|
||||
<tr><td>MMLU-Pro</td><td>acc</td><td>5</td><td>36.54</td><td>29.87</td><td>37.08</td></tr>
|
||||
<tr><td>BBH</td><td>acc</td><td>3</td><td>54.23</td><td>45.58</td><td>53.60</td></tr>
|
||||
<tr><td>SimpleQA<sup>†</sup></td><td>acc</td><td>5</td><td>27.39</td><td>23.76</td><td>16.83</td></tr>
|
||||
<tr><td align="center" colspan="6"><b>Mathematics Tasks</b></td></tr>
|
||||
<tr><td>MATH</td><td>em</td><td>4</td><td>35.08</td><td>30.82</td><td>41.74</td></tr>
|
||||
<tr><td>GSM8K</td><td>em</td><td>8</td><td>61.03</td><td>51.93</td><td>75.74</td></tr>
|
||||
<tr><td align="center" colspan="6"><b>Coding Tasks</b></td></tr>
|
||||
<tr><td>HumanEval</td><td>pass@1</td><td>0</td><td>55.96</td><td>51.33</td><td>45.31</td></tr>
|
||||
<tr><td>MBPP</td><td>pass@1</td><td>3</td><td>50.95</td><td>44.86</td><td>53.85</td></tr>
|
||||
<tr><td align="center" colspan="6"><b>Korean Tasks</b></td></tr>
|
||||
<tr><td>KMMLU</td><td>acc</td><td>5</td><td>47.92</td><td>44.17</td><td>43.08</td></tr>
|
||||
<tr><td>KoSimpleQA<sup>†</sup></td><td>acc</td><td>5</td><td>32.50</td><td>28.50</td><td>13.40</td></tr>
|
||||
<tr><td>HAE-RAE Bench (v1.0)</td><td>acc</td><td>5</td><td>80.66</td><td>75.34</td><td>55.54</td></tr>
|
||||
<tr><td>MATH-Ko<sup>‡</sup></td><td>em</td><td>4</td><td>31.54</td><td>24.75</td><td>31.92</td></tr>
|
||||
<tr><td>MBPP-Ko<sup>§</sup></td><td>pass@1</td><td>3</td><td>44.55</td><td>39.38</td><td>47.97</td></tr>
|
||||
<tr><td align="center" colspan="6"><b>Long Context Tasks</b></td></tr>
|
||||
<tr><td>RULER-32K</td><td>acc</td><td>0</td><td>67.50</td><td>55.72</td><td>69.01</td></tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
|
||||
† Evaluated in Multiple Choice Question Answering (MCQA) format with 10 options.
|
||||
|
||||
‡ Subsets from [HRM8K](https://huggingface.co/datasets/HAERAE-HUB/HRM8K) (MATH, GSM8K).
|
||||
|
||||
§ Internally translated to Korean.
|
||||
|
||||
### Instruct model evaluation results
|
||||
|
||||
Instruction-following, chat, tool-calling, code, math, and knowledge benchmarks for the Kanana-2 SLM series. Scores use greedy decoding (temperature 0.0, top-p 1.0, max 4096 tokens); the metric for each benchmark is listed in the Metric column.
|
||||
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>Benchmark</th>
|
||||
<th>Metric</th>
|
||||
<th>kanana-2-3b-instruct</th>
|
||||
<th>kanana-2-1.3b-instruct</th>
|
||||
<th>Qwen3.5-2B</th>
|
||||
<th>Qwen3-1.7B</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr><td align="center" colspan="6"><b>Chat</b></td></tr>
|
||||
<tr><td>MT-Bench<sup>†</sup></td><td>judge</td><td>7.15</td><td>6.83</td><td>6.87</td><td>6.98</td></tr>
|
||||
<tr><td>KoMT-Bench<sup>†</sup></td><td>judge</td><td>6.92</td><td>6.54</td><td>5.21</td><td>5.29</td></tr>
|
||||
<tr><td align="center" colspan="6"><b>Instruction Following</b></td></tr>
|
||||
<tr><td>IFBench</td><td>prompt strict</td><td>33.33</td><td>34.69</td><td>24.83</td><td>18.33</td></tr>
|
||||
<tr><td>IFEval</td><td>prompt strict</td><td>80.96</td><td>77.63</td><td>66.91</td><td>68.39</td></tr>
|
||||
<tr><td>IHEval</td><td>pass@1</td><td>35.96</td><td>27.06</td><td>38.52</td><td>42.09</td></tr>
|
||||
<tr><td align="center" colspan="6"><b>Tool Calling</b></td></tr>
|
||||
<tr><td>BFCL-v3 (Live)<sup>‡</sup></td><td>pass@1</td><td>71.94</td><td>69.64</td><td>66.96</td><td>65.48</td></tr>
|
||||
<tr><td>BFCL-v3 (Multi-Turn)<sup>‡</sup></td><td>pass@1</td><td>17.12</td><td>5.50</td><td>6.27</td><td>4.38</td></tr>
|
||||
<tr><td align="center" colspan="6"><b>Code Generation</b></td></tr>
|
||||
<tr><td>MBPP</td><td>pass@1</td><td>70.63</td><td>69.05</td><td>55.56</td><td>62.17</td></tr>
|
||||
<tr><td>MBPP+</td><td>pass@1</td><td>60.05</td><td>60.85</td><td>46.83</td><td>52.65</td></tr>
|
||||
<tr><td align="center" colspan="6"><b>Mathematics</b></td></tr>
|
||||
<tr><td>GSM-Plus</td><td>pass@1</td><td>61.17</td><td>58.27</td><td>61.29</td><td>63.10</td></tr>
|
||||
<tr><td>MATH-500</td><td>pass@1</td><td>61.20</td><td>61.40</td><td>67.80</td><td>72.00</td></tr>
|
||||
<tr><td>Minerva Math</td><td>pass@1</td><td>24.44</td><td>22.43</td><td>35.18</td><td>27.21</td></tr>
|
||||
<tr><td align="center" colspan="6"><b>Reasoning & Knowledge</b></td></tr>
|
||||
<tr><td>MMLU-CoT</td><td>acc</td><td>61.09</td><td>60.37</td><td>69.01</td><td>66.02</td></tr>
|
||||
<tr><td>KMMLU-CoT</td><td>acc</td><td>43.32</td><td>42.79</td><td>41.75</td><td>37.84</td></tr>
|
||||
<tr><td>HAERAE-Bench (v1.0)-CoT</td><td>acc</td><td>43.75</td><td>44.89</td><td>27.84</td><td>27.84</td></tr>
|
||||
<tr><td>KoSimpleQA</td><td>acc</td><td>22.29</td><td>17.81</td><td>3.21</td><td>2.82</td></tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
|
||||
† Evaluated using `gpt-4o-2024-08-06` as the judge model.
|
||||
|
||||
‡ `Live` denotes the average score of 6 live benchmarks, and `Multi-Turn` the average score of 4 multi-turn benchmarks.
|
||||
|
||||
## Deployment
|
||||
|
||||
`kanana-2-1.3b-instruct` uses a custom hybrid attention architecture (`Kanana2TinyForCausalLM` — a Qwen3 backbone with a 3:1 SWA/full-attention layout and per-layer-type RoPE), shipped as remote code in the repository.
|
||||
|
||||
> [!NOTE]
|
||||
> Because the modeling code is loaded from the repository, serving requires `transformers >= 4.57` and the `--trust-remote-code` flag. For SGLang, use `--attention-backend triton` so the hybrid sliding-window attention is handled correctly.
|
||||
|
||||
|
||||
|
||||
### vLLM
|
||||
|
||||
[vLLM](https://github.com/vllm-project/vllm) is a fast and memory-optimized engine designed for high-performance LLM inference and serving.
|
||||
|
||||
```shell
|
||||
vllm serve kakaocorp/kanana-2-1.3b-instruct \
|
||||
--tensor-parallel-size 1 \
|
||||
--max-model-len 32768 \
|
||||
--trust-remote-code \
|
||||
--enable-auto-tool-choice \
|
||||
--tool-call-parser qwen3_coder
|
||||
```
|
||||
|
||||
|
||||
|
||||
### SGLang
|
||||
|
||||
[SGLang](https://github.com/sgl-project/sglang) is a high-efficiency framework for serving LLMs and VLMs, enabling easy deployment of OpenAI-compatible API servers.
|
||||
|
||||
For SGLang, the model is served through the stock `Qwen3ForCausalLM` path instead of the remote-code `Kanana2TinyForCausalLM` class. This requires two files shipped in the `sglang/` directory of this repository:
|
||||
|
||||
1. `sglang/config.json` — a Qwen3-flavored config (`architectures: ["Qwen3ForCausalLM"]`, `model_type: qwen3`, no `auto_map`) that keeps the hybrid-attention fields (`layer_types`, `sliding_window`, per-layer-type `rope_parameters`). Use it in place of the default `config.json` when serving with SGLang.
|
||||
2. `sglang/qwen3.py` — a patched model definition that overrides the installed `sglang/srt/models/qwen3.py`.
|
||||
|
||||
```shell
|
||||
python3 -m sglang.launch_server \
|
||||
--model-path kakaocorp/kanana-2-1.3b-instruct \
|
||||
--tp 1 \
|
||||
--context-length 32768 \
|
||||
--attention-backend triton \
|
||||
--trust-remote-code \
|
||||
--tool-call-parser qwen3_coder
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
>
|
||||
> - Recommended: `sglang==0.5.1`.
|
||||
> - With the Qwen3-flavored `sglang/config.json` the model itself needs no remote code, but `--trust-remote-code` is passed so the tokenizer/config are loaded without prompting.
|
||||
> - Use `triton` or `fa3` for the attention backend. **Avoid** `flashinfer` — it appears to have an issue with this model and causes significant accuracy/throughput degradation.
|
||||
|
||||
Kanana-2-1.3B shares the Qwen3 backbone but adds a **3:1 SWA/full hybrid attention layout** and **per-layer-type RoPE**, which SGLang's stock Qwen3 model does not handle. The patched `sglang/qwen3.py` changes the decoder layer to:
|
||||
|
||||
1. **Per-layer-type RoPE** — for each layer, read the entry in `config.rope_parameters` matching `config.layer_types[layer_id]` and build that layer's own rotary embedding from it: `full_attention` uses YaRN (`factor=40`, `original_max_position_embeddings=4096`) and `sliding_attention` uses default RoPE (`rope_theta=10000`). A layer type with no matching entry falls back to no RoPE.
|
||||
2. **Hybrid sliding-window attention** — layers typed `sliding_attention` run `RadixAttention` with `sliding_window_size = config.sliding_window - 1` (SGLang uses an exclusive window, HF an inclusive one), while `full_attention` layers use full causal attention.
|
||||
|
||||
|
||||
|
||||
## License
|
||||
|
||||
The model weights are released under the [KananaOpenLicense](https://huggingface.co/kakaocorp/kanana-2-1.3b-instruct/blob/main/LICENSE).
|
||||
|
||||
## Citation
|
||||
|
||||
```
|
||||
@misc{kanana2slm2026,
|
||||
title = {Kanana-2 SLM},
|
||||
author = {Kanana LLM},
|
||||
year = {2026},
|
||||
url = {https://huggingface.co/collections/kakaocorp/kanana-2-slm}
|
||||
}
|
||||
```
|
||||
|
||||
|
||||
|
||||
## Contact
|
||||
|
||||
- Kanana LLM Team Technical Support: [kanana-llm@kakaocorp.com](mailto:kanana-llm@kakaocorp.com)
|
||||
- Business & Partnership Contact: [alpha.k@kakaocorp.com](mailto:alpha.k@kakaocorp.com)
|
||||
|
||||
3
assets/logo/kanana.png
Normal file
3
assets/logo/kanana.png
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:52b7b1de0b928150c7b8fe3517a7b86b4116271f2f0f0108bf2520281d6bac52
|
||||
size 109633
|
||||
211
chat_template.jinja
Normal file
211
chat_template.jinja
Normal file
@@ -0,0 +1,211 @@
|
||||
{%- if not add_generation_prompt is defined %}
|
||||
{%- set add_generation_prompt = false %}
|
||||
{%- endif %}
|
||||
{%- set thinking_mode = thinking_mode | default('no_think') | lower -%}
|
||||
{%- set image_count = namespace(value=0) %}
|
||||
{%- set video_count = namespace(value=0) %}
|
||||
{%- set audio_count = namespace(value=0) %}
|
||||
{%- macro render_content(content, do_vision_count, is_system_content=false, do_audio_count=false) %}
|
||||
{%- if content is string %}
|
||||
{{- content }}
|
||||
{%- elif content is iterable and content is not mapping %}
|
||||
{%- for item in content %}
|
||||
{%- if 'image' in item or 'image_url' in item or item.type == 'image' %}
|
||||
{%- if is_system_content %}
|
||||
{{- raise_exception('System message cannot contain images.') }}
|
||||
{%- endif %}
|
||||
{%- if do_vision_count %}
|
||||
{%- set image_count.value = image_count.value + 1 %}
|
||||
{%- endif %}
|
||||
{%- if add_vision_id %}
|
||||
{{- 'Picture ' ~ image_count.value ~ ': ' }}
|
||||
{%- endif %}
|
||||
{{- '<|vision_start|><|image_pad|><|vision_end|>' }}
|
||||
{%- elif 'video' in item or item.type == 'video' %}
|
||||
{%- if is_system_content %}
|
||||
{{- raise_exception('System message cannot contain videos.') }}
|
||||
{%- endif %}
|
||||
{%- if do_vision_count %}
|
||||
{%- set video_count.value = video_count.value + 1 %}
|
||||
{%- endif %}
|
||||
{%- if add_vision_id %}
|
||||
{{- 'Video ' ~ video_count.value ~ ': ' }}
|
||||
{%- endif %}
|
||||
{{- '<|vision_start|><|video_pad|><|vision_end|>' }}
|
||||
{%- elif 'audio' in item or item.type == 'audio' %}
|
||||
{%- if is_system_content %}
|
||||
{{- raise_exception('System message cannot contain audios.') }}
|
||||
{%- endif %}
|
||||
{%- if do_audio_count %}
|
||||
{%- set audio_count.value = audio_count.value + 1 %}
|
||||
{%- endif %}
|
||||
{%- if add_audio_id %}
|
||||
{{- 'Audio ' ~ audio_count.value ~ ': ' }}
|
||||
{%- endif %}
|
||||
{{- '<|audio_start|><|audio_pad|><|audio_end|>' }}
|
||||
{%- elif 'text' in item %}
|
||||
{{- item.text }}
|
||||
{%- else %}
|
||||
{{- raise_exception('Unexpected item type in content.') }}
|
||||
{%- endif %}
|
||||
{%- endfor %}
|
||||
{%- elif content is none or content is undefined %}
|
||||
{{- '' }}
|
||||
{%- else %}
|
||||
{{- raise_exception('Unexpected content type.') }}
|
||||
{%- endif %}
|
||||
{%- endmacro %}
|
||||
{%- if not messages %}
|
||||
{{- raise_exception('No messages provided.') }}
|
||||
{%- endif %}
|
||||
{%- if tools and tools is iterable and tools is not mapping %}
|
||||
{{- '<|im_start|>system\n' }}
|
||||
{{- "# Tools\n\nYou have access to the following functions:\n\n<tools>" }}
|
||||
{%- for tool in tools %}
|
||||
{{- "\n" }}
|
||||
{{- tool | tojson }}
|
||||
{%- endfor %}
|
||||
{{- "\n</tools>" }}
|
||||
{{- '\n\nIf you choose to call a function ONLY reply in the following format with NO suffix:\n\n<tool_call>\n<function=example_function_name>\n<parameter=example_parameter_1>\nvalue_1\n</parameter>\n<parameter=example_parameter_2>\nThis is the value for the second parameter\nthat can span\nmultiple lines\n</parameter>\n</function>\n</tool_call>\n\n<IMPORTANT>\nReminder:\n- Function calls MUST follow the specified format: an inner <function=...></function> block must be nested within <tool_call></tool_call> XML tags\n- Required parameters MUST be specified\n- You may provide optional reasoning for your function call in natural language BEFORE the function call, but NOT after\n- If there is no function call available, answer the question like normal with your current knowledge and do not tell the user about function calls\n</IMPORTANT>' }}
|
||||
{%- if messages[0].role == 'system' %}
|
||||
{%- set content = render_content(messages[0].content, false, true)|trim %}
|
||||
{%- if content %}
|
||||
{{- '\n\n' + content }}
|
||||
{%- endif %}
|
||||
{%- endif %}
|
||||
{{- '<|im_end|>\n' }}
|
||||
{%- else %}
|
||||
{%- if messages[0].role == 'system' %}
|
||||
{%- set content = render_content(messages[0].content, false, true)|trim %}
|
||||
{{- '<|im_start|>system\n' + content + '<|im_end|>\n' }}
|
||||
{%- endif %}
|
||||
{%- endif %}
|
||||
{%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %}
|
||||
{%- for message in messages[::-1] %}
|
||||
{%- set index = (messages|length - 1) - loop.index0 %}
|
||||
{%- if ns.multi_step_tool and message.role == "user" %}
|
||||
{%- set content = render_content(message.content, false)|trim %}
|
||||
{%- if not(content.startswith('<tool_response>') and content.endswith('</tool_response>')) %}
|
||||
{%- set ns.multi_step_tool = false %}
|
||||
{%- set ns.last_query_index = index %}
|
||||
{%- endif %}
|
||||
{%- endif %}
|
||||
{%- endfor %}
|
||||
{%- if ns.multi_step_tool %}
|
||||
{{- raise_exception('No user query found in messages.') }}
|
||||
{%- endif %}
|
||||
{%- for message in messages %}
|
||||
{%- set content = render_content(message.content, true, do_audio_count=do_audio_count)|trim %}
|
||||
{%- if message.role == "system" %}
|
||||
{%- if not loop.first %}
|
||||
{{- raise_exception('System message must be at the beginning.') }}
|
||||
{%- endif %}
|
||||
{%- elif message.role == "user" %}
|
||||
{{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }}
|
||||
{%- elif message.role == "assistant" %}
|
||||
{%- set reasoning_content = '' %}
|
||||
{%- if message.reasoning_content is string %}
|
||||
{%- set reasoning_content = message.reasoning_content %}
|
||||
{%- else %}
|
||||
{%- if '</think>' in content %}
|
||||
{%- set reasoning_content = content.split('</think>')[0].rstrip('\n').split('<think>')[-1].lstrip('\n') %}
|
||||
{%- set content = content.split('</think>')[-1].lstrip('\n') %}
|
||||
{%- endif %}
|
||||
{%- endif %}
|
||||
{%- set reasoning_content = reasoning_content|trim %}
|
||||
{%- if (preserve_thinking is defined and preserve_thinking is true) or (loop.index0 > ns.last_query_index) %}
|
||||
{#- prefix (outside generation, no loss) -#}
|
||||
{%- if reasoning_content %}
|
||||
{{- '<|im_start|>' + message.role + '\n<think>\n' -}}
|
||||
{%- else %}
|
||||
{{- '<|im_start|>' + message.role + '\n<think>\n\n</think>\n\n' -}}
|
||||
{%- endif %}
|
||||
{%- generation -%}
|
||||
{#- body (inside generation, loss flows here) -#}
|
||||
{%- if reasoning_content %}
|
||||
{{- reasoning_content + '\n</think>\n\n' + content }}
|
||||
{%- else %}
|
||||
{{- content }}
|
||||
{%- endif %}
|
||||
{%- if message.tool_calls and message.tool_calls is iterable and message.tool_calls is not mapping %}
|
||||
{%- for tool_call in message.tool_calls %}
|
||||
{%- if tool_call.function is defined %}
|
||||
{%- set tool_call = tool_call.function %}
|
||||
{%- endif %}
|
||||
{%- if loop.first %}
|
||||
{%- if content|trim %}
|
||||
{{- '\n\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
|
||||
{%- else %}
|
||||
{{- '<tool_call>\n<function=' + tool_call.name + '>\n' }}
|
||||
{%- endif %}
|
||||
{%- else %}
|
||||
{{- '\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
|
||||
{%- endif %}
|
||||
{%- if tool_call.arguments is mapping %}
|
||||
{%- for args_name, args_value in tool_call.arguments|items %}
|
||||
{{- '<parameter=' + args_name + '>\n' }}
|
||||
{%- set args_value = args_value | string if args_value is string else args_value | tojson | safe %}
|
||||
{{- args_value }}
|
||||
{{- '\n</parameter>\n' }}
|
||||
{%- endfor %}
|
||||
{%- endif %}
|
||||
{{- '</function>\n</tool_call>' }}
|
||||
{%- endfor %}
|
||||
{%- endif %}
|
||||
{{- '<|im_end|>\n' -}}
|
||||
{%- endgeneration -%}
|
||||
{%- else %}
|
||||
{{- '<|im_start|>' + message.role + '\n' + content }}
|
||||
{%- if message.tool_calls and message.tool_calls is iterable and message.tool_calls is not mapping %}
|
||||
{%- for tool_call in message.tool_calls %}
|
||||
{%- if tool_call.function is defined %}
|
||||
{%- set tool_call = tool_call.function %}
|
||||
{%- endif %}
|
||||
{%- if loop.first %}
|
||||
{%- if content|trim %}
|
||||
{{- '\n\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
|
||||
{%- else %}
|
||||
{{- '<tool_call>\n<function=' + tool_call.name + '>\n' }}
|
||||
{%- endif %}
|
||||
{%- else %}
|
||||
{{- '\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
|
||||
{%- endif %}
|
||||
{%- if tool_call.arguments is mapping %}
|
||||
{%- for args_name, args_value in tool_call.arguments|items %}
|
||||
{{- '<parameter=' + args_name + '>\n' }}
|
||||
{%- set args_value = args_value | string if args_value is string else args_value | tojson | safe %}
|
||||
{{- args_value }}
|
||||
{{- '\n</parameter>\n' }}
|
||||
{%- endfor %}
|
||||
{%- endif %}
|
||||
{{- '</function>\n</tool_call>' }}
|
||||
{%- endfor %}
|
||||
{%- endif %}
|
||||
{{- '<|im_end|>\n' }}
|
||||
{%- endif %}
|
||||
{%- elif message.role == "tool" %}
|
||||
{%- if loop.previtem and loop.previtem.role != "tool" %}
|
||||
{{- '<|im_start|>user' }}
|
||||
{%- endif %}
|
||||
{{- '\n<tool_response>\n' }}
|
||||
{{- content }}
|
||||
{{- '\n</tool_response>' }}
|
||||
{%- if not loop.last and loop.nextitem.role != "tool" %}
|
||||
{{- '<|im_end|>\n' }}
|
||||
{%- elif loop.last %}
|
||||
{{- '<|im_end|>\n' }}
|
||||
{%- endif %}
|
||||
{%- else %}
|
||||
{{- raise_exception('Unexpected message role.') }}
|
||||
{%- endif %}
|
||||
{%- endfor %}
|
||||
{%- if add_generation_prompt %}
|
||||
{{- '<|im_start|>assistant\n' }}
|
||||
{%- if thinking_mode == 'no_think' %}
|
||||
{{- '<think>\n\n</think>\n\n' -}}
|
||||
{%- elif thinking_mode == 'auto' %}
|
||||
{{- '<think>' -}}
|
||||
{%- else %}
|
||||
{{- '<think>\n' -}}
|
||||
{%- endif %}
|
||||
{%- endif %}
|
||||
81
config.json
Normal file
81
config.json
Normal file
@@ -0,0 +1,81 @@
|
||||
{
|
||||
"architectures": [
|
||||
"Kanana2TinyForCausalLM"
|
||||
],
|
||||
"attention_bias": false,
|
||||
"attention_dropout": 0.0,
|
||||
"auto_map": {
|
||||
"AutoConfig": "configuration_kanana2_tiny.Kanana2TinyConfig",
|
||||
"AutoModel": "modeling_kanana2_tiny.Kanana2TinyModel",
|
||||
"AutoModelForCausalLM": "modeling_kanana2_tiny.Kanana2TinyForCausalLM"
|
||||
},
|
||||
"bos_token_id": 128000,
|
||||
"dtype": "bfloat16",
|
||||
"eos_token_id": 128010,
|
||||
"head_dim": 128,
|
||||
"hidden_act": "silu",
|
||||
"hidden_size": 1280,
|
||||
"initializer_range": 0.02,
|
||||
"intermediate_size": 5760,
|
||||
"layer_types": [
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention"
|
||||
],
|
||||
"max_position_embeddings": 32768,
|
||||
"max_window_layers": 32,
|
||||
"model_type": "kanana2_tiny",
|
||||
"num_attention_heads": 32,
|
||||
"num_hidden_layers": 32,
|
||||
"num_key_value_heads": 8,
|
||||
"pad_token_id": 128001,
|
||||
"rms_norm_eps": 1e-06,
|
||||
"rope_parameters": {
|
||||
"full_attention": {
|
||||
"factor": 40.0,
|
||||
"original_max_position_embeddings": 4096,
|
||||
"rope_theta": 10000,
|
||||
"rope_type": "yarn"
|
||||
},
|
||||
"sliding_attention": {
|
||||
"rope_theta": 10000.0,
|
||||
"rope_type": "default"
|
||||
}
|
||||
},
|
||||
"rope_scaling": null,
|
||||
"rope_theta": 10000,
|
||||
"sliding_window": 1024,
|
||||
"transformers_version": "4.57.1",
|
||||
"use_cache": true,
|
||||
"use_sliding_window": true,
|
||||
"vocab_size": 128256
|
||||
}
|
||||
160
configuration_kanana2_tiny.py
Normal file
160
configuration_kanana2_tiny.py
Normal file
@@ -0,0 +1,160 @@
|
||||
# coding=utf-8
|
||||
# Configuration class for Kanana-2 PD-series (Qwen3 architecture with
|
||||
# sliding/full alternating attention and per-attention-type RoPE).
|
||||
#
|
||||
# Difference vs Qwen3:
|
||||
# * `rope_parameters` is a dict keyed by attention type (`full_attention` /
|
||||
# `sliding_attention`). Each entry is a self-contained RoPE config
|
||||
# understood by `transformers.modeling_rope_utils.ROPE_INIT_FUNCTIONS`.
|
||||
# This lets us apply YaRN to global-attention layers while keeping
|
||||
# unscaled RoPE for sliding-attention layers.
|
||||
# * Top-level `rope_scaling` is unused on this config; the modeling code
|
||||
# builds per-attention-type sub-configs at construction time and sets
|
||||
# `rope_scaling` on each sub-config so HF's standard rope init functions
|
||||
# (which read `config.rope_scaling`) work unchanged.
|
||||
|
||||
from transformers.configuration_utils import PretrainedConfig, layer_type_validation
|
||||
from transformers.modeling_rope_utils import rope_config_validation
|
||||
from transformers.utils import logging
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
|
||||
class Kanana2TinyConfig(PretrainedConfig):
|
||||
"""Configuration for the Kanana-2 PD-series (Qwen3 + per-type RoPE)."""
|
||||
|
||||
model_type = "kanana2_tiny"
|
||||
keys_to_ignore_at_inference = ["past_key_values"]
|
||||
|
||||
base_model_tp_plan = {
|
||||
"layers.*.self_attn.q_proj": "colwise",
|
||||
"layers.*.self_attn.k_proj": "colwise",
|
||||
"layers.*.self_attn.v_proj": "colwise",
|
||||
"layers.*.self_attn.o_proj": "rowwise",
|
||||
"layers.*.mlp.gate_proj": "colwise",
|
||||
"layers.*.mlp.up_proj": "colwise",
|
||||
"layers.*.mlp.down_proj": "rowwise",
|
||||
}
|
||||
base_model_pp_plan = {
|
||||
"embed_tokens": (["input_ids"], ["inputs_embeds"]),
|
||||
"layers": (["hidden_states", "attention_mask"], ["hidden_states"]),
|
||||
"norm": (["hidden_states"], ["hidden_states"]),
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size=128256,
|
||||
hidden_size=1024,
|
||||
intermediate_size=4608,
|
||||
num_hidden_layers=32,
|
||||
num_attention_heads=32,
|
||||
num_key_value_heads=8,
|
||||
head_dim=128,
|
||||
hidden_act="silu",
|
||||
max_position_embeddings=35000,
|
||||
initializer_range=0.02,
|
||||
rms_norm_eps=1e-6,
|
||||
use_cache=True,
|
||||
tie_word_embeddings=True,
|
||||
rope_theta=10000.0,
|
||||
rope_parameters=None,
|
||||
rope_scaling=None,
|
||||
attention_bias=False,
|
||||
use_sliding_window=True,
|
||||
sliding_window=1024,
|
||||
max_window_layers=32,
|
||||
layer_types=None,
|
||||
attention_dropout=0.0,
|
||||
**kwargs,
|
||||
):
|
||||
# Standard Qwen3-ish fields
|
||||
self.vocab_size = vocab_size
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
self.hidden_size = hidden_size
|
||||
self.intermediate_size = intermediate_size
|
||||
self.num_hidden_layers = num_hidden_layers
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.use_sliding_window = use_sliding_window
|
||||
self.sliding_window = sliding_window if self.use_sliding_window else None
|
||||
self.max_window_layers = max_window_layers
|
||||
|
||||
if num_key_value_heads is None:
|
||||
num_key_value_heads = num_attention_heads
|
||||
self.num_key_value_heads = num_key_value_heads
|
||||
self.head_dim = head_dim
|
||||
self.hidden_act = hidden_act
|
||||
self.initializer_range = initializer_range
|
||||
self.rms_norm_eps = rms_norm_eps
|
||||
self.use_cache = use_cache
|
||||
self.rope_theta = rope_theta
|
||||
self.attention_bias = attention_bias
|
||||
self.attention_dropout = attention_dropout
|
||||
# Kept for HF helpers that probe the attribute. The per-attention RoPE
|
||||
# config lives in `rope_parameters`; modeling code constructs sub-configs
|
||||
# whose `rope_scaling` is the per-type dict at init time.
|
||||
self.rope_scaling = rope_scaling
|
||||
|
||||
# Per-attention-type RoPE.
|
||||
# Expected shape (defaults match the kanana-2-pd-series checkpoints):
|
||||
# {
|
||||
# "full_attention": {"rope_type": "yarn", "rope_theta": 10000,
|
||||
# "factor": 40.0, "original_max_position_embeddings": 4096},
|
||||
# "sliding_attention": {"rope_type": "default", "rope_theta": 10000.0},
|
||||
# }
|
||||
if rope_parameters is None:
|
||||
rope_parameters = {
|
||||
"full_attention": {
|
||||
"rope_type": "default",
|
||||
"rope_theta": rope_theta,
|
||||
},
|
||||
"sliding_attention": {
|
||||
"rope_type": "default",
|
||||
"rope_theta": rope_theta,
|
||||
},
|
||||
}
|
||||
self.rope_parameters = rope_parameters
|
||||
|
||||
for attn_type, params in self.rope_parameters.items():
|
||||
if not isinstance(params, dict) or "rope_type" not in params:
|
||||
raise ValueError(
|
||||
f"rope_parameters[{attn_type!r}] must be a dict with a 'rope_type' key, got {params!r}"
|
||||
)
|
||||
|
||||
# Set layer_types BEFORE per-type rope validation: the layer types must
|
||||
# exist for the validators that gate on layer_types.
|
||||
self.layer_types = layer_types
|
||||
if self.layer_types is None:
|
||||
self.layer_types = [
|
||||
"sliding_attention"
|
||||
if self.sliding_window is not None and i >= self.max_window_layers
|
||||
else "full_attention"
|
||||
for i in range(self.num_hidden_layers)
|
||||
]
|
||||
layer_type_validation(self.layer_types, self.num_hidden_layers)
|
||||
|
||||
# Per-attention-type rope validation. The 4.57.1 validators read off
|
||||
# `config.rope_scaling` (flat dict) and `config.rope_theta` (top-level),
|
||||
# so for each per-type sub-dict we present it in that shape, run the
|
||||
# validator, then restore. `rope_theta` is filtered out of the temporary
|
||||
# `rope_scaling` because in 4.57.1's schema it lives at the top level.
|
||||
for attn_type, params in self.rope_parameters.items():
|
||||
if attn_type not in set(self.layer_types):
|
||||
continue
|
||||
saved_rope_scaling = self.rope_scaling
|
||||
saved_rope_theta = self.rope_theta
|
||||
try:
|
||||
self.rope_scaling = {k: v for k, v in params.items() if k != "rope_theta"}
|
||||
self.rope_theta = params.get("rope_theta", saved_rope_theta)
|
||||
rope_config_validation(self)
|
||||
finally:
|
||||
self.rope_scaling = saved_rope_scaling
|
||||
self.rope_theta = saved_rope_theta
|
||||
|
||||
super().__init__(
|
||||
tie_word_embeddings=tie_word_embeddings,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["Kanana2TinyConfig"]
|
||||
7
generation_config.json
Normal file
7
generation_config.json
Normal file
@@ -0,0 +1,7 @@
|
||||
{
|
||||
"_from_model_config": true,
|
||||
"bos_token_id": 128000,
|
||||
"eos_token_id": 128010,
|
||||
"pad_token_id": 128001,
|
||||
"transformers_version": "4.57.1"
|
||||
}
|
||||
3
model-00001-of-00002.safetensors
Normal file
3
model-00001-of-00002.safetensors
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:fd6e10573d4f630037806d3104ee2fe789a909063c38d9983412556beac8f020
|
||||
size 4995531560
|
||||
3
model-00002-of-00002.safetensors
Normal file
3
model-00002-of-00002.safetensors
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:c6bbbd8f06f39b73f17dcdd1e2342e1b3f0405d446ee6c9df86ff6834629734f
|
||||
size 827092704
|
||||
3
model.safetensors
Normal file
3
model.safetensors
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:49aa6cd8686563c59321d83810731956c61ec8d5c8538a249d38007986cdc942
|
||||
size 2582997160
|
||||
363
model.safetensors.index.json
Normal file
363
model.safetensors.index.json
Normal file
@@ -0,0 +1,363 @@
|
||||
{
|
||||
"metadata": {
|
||||
"total_parameters": 1455645952,
|
||||
"total_size": 5822583808
|
||||
},
|
||||
"weight_map": {
|
||||
"lm_head.weight": "model-00002-of-00002.safetensors",
|
||||
"model.embed_tokens.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.0.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.0.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.0.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.0.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.0.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.0.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.0.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.0.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.0.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.0.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.0.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.1.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.1.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.1.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.1.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.1.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.1.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.1.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.1.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.1.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.1.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.1.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.10.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.10.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.10.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.10.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.10.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.10.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.10.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.10.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.10.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.10.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.10.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.11.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.11.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.11.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.11.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.11.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.11.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.11.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.11.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.11.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.11.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.11.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.12.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.12.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.12.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.12.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.12.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.12.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.12.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.12.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.12.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.12.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.12.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.13.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.13.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.13.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.13.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.13.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.13.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.13.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.13.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.13.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.13.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.13.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.14.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.14.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.14.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.14.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.14.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.14.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.14.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.14.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.14.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.14.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.14.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.15.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.15.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.15.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.15.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.15.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.15.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.15.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.15.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.15.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.15.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.15.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.16.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.16.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.16.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.16.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.16.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.16.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.16.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.16.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.16.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.16.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.16.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.17.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.17.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.17.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.17.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.17.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.17.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.17.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.17.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.17.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.17.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.17.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.18.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.18.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.18.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.18.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.18.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.18.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.18.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.18.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.18.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.18.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.18.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.19.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.19.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.19.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.19.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.19.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.19.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.19.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.19.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.19.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.19.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.19.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.2.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.2.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.2.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.2.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.2.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.2.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.2.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.2.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.2.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.2.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.2.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.20.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.20.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.20.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.20.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.20.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.20.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.20.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.20.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.20.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.20.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.20.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.21.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.21.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.21.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.21.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.21.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.21.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.21.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.21.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.21.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.21.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.21.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.22.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.22.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.22.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.22.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.22.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.22.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.22.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.22.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.22.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.22.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.22.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.23.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.23.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.23.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.23.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.23.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.23.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.23.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.23.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.23.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.23.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.23.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.24.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.24.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.24.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.24.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.24.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.24.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.24.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.24.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.24.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.24.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.24.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.25.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.25.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.25.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.25.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.25.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.25.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.25.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.25.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.25.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.25.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.25.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.26.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.26.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.26.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.26.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.26.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.26.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.26.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.26.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.26.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.26.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.26.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.27.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.27.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.27.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.27.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.27.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.27.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.27.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.27.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.27.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.27.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.27.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.28.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.28.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.28.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.28.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.28.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.28.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.28.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.28.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.28.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.28.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.28.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.29.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.29.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.29.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.29.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.29.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.29.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.29.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.29.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.29.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.29.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.29.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.3.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.3.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.3.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.3.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.3.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.3.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.3.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.3.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.3.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.3.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.3.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.30.input_layernorm.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.30.mlp.down_proj.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.30.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.30.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.30.post_attention_layernorm.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.30.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.30.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.30.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.30.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.30.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.30.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.31.input_layernorm.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.31.mlp.down_proj.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.31.mlp.gate_proj.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.31.mlp.up_proj.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.31.post_attention_layernorm.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.31.self_attn.k_norm.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.31.self_attn.k_proj.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.31.self_attn.o_proj.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.31.self_attn.q_norm.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.31.self_attn.q_proj.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.31.self_attn.v_proj.weight": "model-00002-of-00002.safetensors",
|
||||
"model.layers.4.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.4.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.4.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.4.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.4.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.4.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.4.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.4.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.4.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.4.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.4.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.5.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.5.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.5.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.5.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.5.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.5.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.5.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.5.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.5.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.5.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.5.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.6.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.6.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.6.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.6.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.6.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.6.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.6.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.6.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.6.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.6.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.6.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.7.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.7.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.7.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.7.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.7.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.7.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.7.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.7.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.7.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.7.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.7.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.8.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.8.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.8.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.8.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.8.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.8.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.8.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.8.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.8.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.8.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.8.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.9.input_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.9.mlp.down_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.9.mlp.gate_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.9.mlp.up_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.9.post_attention_layernorm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.9.self_attn.k_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.9.self_attn.k_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.9.self_attn.o_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.9.self_attn.q_norm.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.9.self_attn.q_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.layers.9.self_attn.v_proj.weight": "model-00001-of-00002.safetensors",
|
||||
"model.norm.weight": "model-00002-of-00002.safetensors"
|
||||
}
|
||||
}
|
||||
625
modeling_kanana2_tiny.py
Normal file
625
modeling_kanana2_tiny.py
Normal file
@@ -0,0 +1,625 @@
|
||||
# coding=utf-8
|
||||
# Modeling code for the Kanana-2 PD-series (Qwen3 backbone with sliding/full
|
||||
# alternating attention and per-attention-type RoPE).
|
||||
#
|
||||
# Implementation strategy
|
||||
# -----------------------
|
||||
# The architecture is identical to Qwen3 except that the rotary embedding
|
||||
# differs between full-attention and sliding-attention layers. We therefore:
|
||||
# * keep the exact Qwen3 layer/attention/MLP/RMSNorm code (copied here so the
|
||||
# module is self-contained for `trust_remote_code=True` loading), and
|
||||
# * instantiate two rotary embeddings — one per attention type — and dispatch
|
||||
# to the right one in each decoder layer based on `layer_types`.
|
||||
#
|
||||
# The trick for "two rotary embeddings driven by one shared config" follows the
|
||||
# Gemma3 pattern: deepcopy the config and overwrite `rope_theta` / `rope_scaling`
|
||||
# to whatever the corresponding `config.rope_parameters[attention_type]` says,
|
||||
# then construct a standard rotary embedding from it.
|
||||
|
||||
import copy
|
||||
from typing import Callable, Optional, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from transformers.activations import ACT2FN
|
||||
from transformers.cache_utils import Cache, DynamicCache
|
||||
from transformers.generation import GenerationMixin
|
||||
from transformers.masking_utils import create_causal_mask, create_sliding_window_causal_mask
|
||||
from transformers.modeling_flash_attention_utils import FlashAttentionKwargs
|
||||
from transformers.modeling_layers import GradientCheckpointingLayer
|
||||
from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast
|
||||
from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update
|
||||
from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel
|
||||
from transformers.processing_utils import Unpack
|
||||
from transformers.utils import TransformersKwargs, auto_docstring, can_return_tuple
|
||||
from transformers.utils.deprecation import deprecate_kwarg
|
||||
|
||||
# ── Cross-version compatibility shims ──────────────────────────────────────
|
||||
# Feature-detection (not version-string compare) because the Kakao-patched
|
||||
# transformers 5.3.0 selectively backports newer APIs, so plain version
|
||||
# inequalities give wrong answers on patched builds.
|
||||
#
|
||||
# Three points of divergence we handle here:
|
||||
#
|
||||
# 1. ``transformers.utils.generic.check_model_inputs`` — added around stock
|
||||
# 5.5; absent on Kakao-patched 5.3. Fall back to a no-op decorator.
|
||||
#
|
||||
# 2. ``create_causal_mask`` / ``create_sliding_window_causal_mask`` kwargs:
|
||||
# - ``input_embeds`` accepted ≤5.5 (deprecation alias); removed ≥5.6
|
||||
# - ``inputs_embeds`` accepted ≥5.3 (patched) / ≥5.5 (stock)
|
||||
# - ``cache_position`` accepted ≤5.8; removed ≥5.9
|
||||
# We pick the right embeds-kwarg name and filter out any kwarg the
|
||||
# installed version doesn't take.
|
||||
#
|
||||
# 3. ``ROPE_INIT_FUNCTIONS`` registry:
|
||||
# - Stock ≥5.5 has ``'proportional'`` (renamed from ``'default'``)
|
||||
# - Kakao-patched 5.3 has neither ``'default'`` nor ``'proportional'``
|
||||
# We supply a local fallback for the unscaled-RoPE init when the
|
||||
# registry is missing both keys.
|
||||
import inspect as _inspect_compat # noqa: E402
|
||||
|
||||
try:
|
||||
from transformers.utils.generic import check_model_inputs # noqa: F401
|
||||
except ImportError:
|
||||
def check_model_inputs(fn): # type: ignore[no-redef]
|
||||
return fn
|
||||
|
||||
_CAUSAL_MASK_PARAMS = set(_inspect_compat.signature(create_causal_mask).parameters)
|
||||
_MASK_EMBEDS_KW = (
|
||||
"inputs_embeds" if "inputs_embeds" in _CAUSAL_MASK_PARAMS else "input_embeds"
|
||||
)
|
||||
|
||||
|
||||
def _filter_mask_kwargs(kwargs: dict) -> dict:
|
||||
"""Drop kwargs the installed ``create_causal_mask`` doesn't accept."""
|
||||
return {k: v for k, v in kwargs.items() if k in _CAUSAL_MASK_PARAMS}
|
||||
|
||||
|
||||
def _compute_default_rope_inv_freq(config, device=None, seq_len=None):
|
||||
"""Unscaled-RoPE inv_freq + attention scaling = 1.0. Mirrors transformers'
|
||||
canonical ``compute_default_rope_parameters`` — used when neither
|
||||
``'default'`` nor ``'proportional'`` is in ``ROPE_INIT_FUNCTIONS``.
|
||||
"""
|
||||
if hasattr(config, "rope_parameters") and isinstance(config.rope_parameters, dict) \
|
||||
and "rope_theta" in config.rope_parameters:
|
||||
base = config.rope_parameters["rope_theta"]
|
||||
else:
|
||||
base = getattr(config, "rope_theta", 10000.0)
|
||||
dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads
|
||||
inv_freq = 1.0 / (
|
||||
base ** (
|
||||
torch.arange(0, dim, 2, dtype=torch.int64).to(device=device, dtype=torch.float) / dim
|
||||
)
|
||||
)
|
||||
return inv_freq, 1.0
|
||||
|
||||
|
||||
def _resolve_rope_init(rope_type: str):
|
||||
"""Pick a rope-init callable for ``rope_type`` across versions."""
|
||||
if rope_type in ROPE_INIT_FUNCTIONS:
|
||||
return ROPE_INIT_FUNCTIONS[rope_type]
|
||||
# 'default' was renamed 'proportional' in stock ≥5.5 — try the other name.
|
||||
if rope_type == "default" and "proportional" in ROPE_INIT_FUNCTIONS:
|
||||
return ROPE_INIT_FUNCTIONS["proportional"]
|
||||
if rope_type == "proportional" and "default" in ROPE_INIT_FUNCTIONS:
|
||||
return ROPE_INIT_FUNCTIONS["default"]
|
||||
if rope_type in ("default", "proportional"):
|
||||
return _compute_default_rope_inv_freq
|
||||
raise KeyError(
|
||||
f"rope_type={rope_type!r} not in ROPE_INIT_FUNCTIONS and no fallback "
|
||||
f"available; keys={sorted(ROPE_INIT_FUNCTIONS)}"
|
||||
)
|
||||
|
||||
|
||||
del _inspect_compat
|
||||
# ───────────────────────────────────────────────────────────────────────────
|
||||
|
||||
from .configuration_kanana2_tiny import Kanana2TinyConfig
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Building blocks (copied verbatim from Qwen3)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class Kanana2TinyRMSNorm(nn.Module):
|
||||
def __init__(self, hidden_size, eps: float = 1e-6) -> None:
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(hidden_size))
|
||||
self.variance_epsilon = eps
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
input_dtype = hidden_states.dtype
|
||||
hidden_states = hidden_states.to(torch.float32)
|
||||
variance = hidden_states.pow(2).mean(-1, keepdim=True)
|
||||
hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
|
||||
return self.weight * hidden_states.to(input_dtype)
|
||||
|
||||
def extra_repr(self):
|
||||
return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}"
|
||||
|
||||
|
||||
class Kanana2TinyMLP(nn.Module):
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.hidden_size = config.hidden_size
|
||||
self.intermediate_size = config.intermediate_size
|
||||
self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
|
||||
self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
|
||||
self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
|
||||
self.act_fn = ACT2FN[config.hidden_act]
|
||||
|
||||
def forward(self, x):
|
||||
return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
|
||||
|
||||
|
||||
def rotate_half(x):
|
||||
x1 = x[..., : x.shape[-1] // 2]
|
||||
x2 = x[..., x.shape[-1] // 2 :]
|
||||
return torch.cat((-x2, x1), dim=-1)
|
||||
|
||||
|
||||
def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1):
|
||||
cos = cos.unsqueeze(unsqueeze_dim)
|
||||
sin = sin.unsqueeze(unsqueeze_dim)
|
||||
q_embed = (q * cos) + (rotate_half(q) * sin)
|
||||
k_embed = (k * cos) + (rotate_half(k) * sin)
|
||||
return q_embed, k_embed
|
||||
|
||||
|
||||
def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
|
||||
batch, num_key_value_heads, slen, head_dim = hidden_states.shape
|
||||
if n_rep == 1:
|
||||
return hidden_states
|
||||
hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
|
||||
return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
|
||||
|
||||
|
||||
def eager_attention_forward(
|
||||
module: nn.Module,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor],
|
||||
scaling: float,
|
||||
dropout: float = 0.0,
|
||||
**kwargs: Unpack[TransformersKwargs],
|
||||
):
|
||||
key_states = repeat_kv(key, module.num_key_value_groups)
|
||||
value_states = repeat_kv(value, module.num_key_value_groups)
|
||||
|
||||
attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling
|
||||
if attention_mask is not None:
|
||||
causal_mask = attention_mask[:, :, :, : key_states.shape[-2]]
|
||||
attn_weights = attn_weights + causal_mask
|
||||
|
||||
attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype)
|
||||
attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training)
|
||||
attn_output = torch.matmul(attn_weights, value_states)
|
||||
attn_output = attn_output.transpose(1, 2).contiguous()
|
||||
return attn_output, attn_weights
|
||||
|
||||
|
||||
class Kanana2TinyAttention(nn.Module):
|
||||
"""Multi-headed attention (identical to Qwen3Attention)."""
|
||||
|
||||
def __init__(self, config: Kanana2TinyConfig, layer_idx: int):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.layer_idx = layer_idx
|
||||
self.head_dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
|
||||
self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads
|
||||
self.scaling = self.head_dim**-0.5
|
||||
self.attention_dropout = config.attention_dropout
|
||||
self.is_causal = True
|
||||
|
||||
self.q_proj = nn.Linear(
|
||||
config.hidden_size, config.num_attention_heads * self.head_dim, bias=config.attention_bias
|
||||
)
|
||||
self.k_proj = nn.Linear(
|
||||
config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias
|
||||
)
|
||||
self.v_proj = nn.Linear(
|
||||
config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias
|
||||
)
|
||||
self.o_proj = nn.Linear(
|
||||
config.num_attention_heads * self.head_dim, config.hidden_size, bias=config.attention_bias
|
||||
)
|
||||
self.q_norm = Kanana2TinyRMSNorm(self.head_dim, eps=config.rms_norm_eps)
|
||||
self.k_norm = Kanana2TinyRMSNorm(self.head_dim, eps=config.rms_norm_eps)
|
||||
self.sliding_window = config.sliding_window if config.layer_types[layer_idx] == "sliding_attention" else None
|
||||
|
||||
@deprecate_kwarg("past_key_value", new_name="past_key_values", version="4.58")
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor],
|
||||
attention_mask: Optional[torch.Tensor],
|
||||
past_key_values: Optional[Cache] = None,
|
||||
cache_position: Optional[torch.LongTensor] = None,
|
||||
**kwargs: Unpack[FlashAttentionKwargs],
|
||||
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||
input_shape = hidden_states.shape[:-1]
|
||||
hidden_shape = (*input_shape, -1, self.head_dim)
|
||||
|
||||
query_states = self.q_norm(self.q_proj(hidden_states).view(hidden_shape)).transpose(1, 2)
|
||||
key_states = self.k_norm(self.k_proj(hidden_states).view(hidden_shape)).transpose(1, 2)
|
||||
value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)
|
||||
|
||||
cos, sin = position_embeddings
|
||||
query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)
|
||||
|
||||
if past_key_values is not None:
|
||||
cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position}
|
||||
key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx, cache_kwargs)
|
||||
|
||||
attention_interface: Callable = eager_attention_forward
|
||||
if self.config._attn_implementation != "eager":
|
||||
attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]
|
||||
|
||||
attn_output, attn_weights = attention_interface(
|
||||
self,
|
||||
query_states,
|
||||
key_states,
|
||||
value_states,
|
||||
attention_mask,
|
||||
dropout=0.0 if not self.training else self.attention_dropout,
|
||||
scaling=self.scaling,
|
||||
sliding_window=self.sliding_window,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
attn_output = attn_output.reshape(*input_shape, -1).contiguous()
|
||||
attn_output = self.o_proj(attn_output)
|
||||
return attn_output, attn_weights
|
||||
|
||||
|
||||
class Kanana2TinyDecoderLayer(GradientCheckpointingLayer):
|
||||
def __init__(self, config: Kanana2TinyConfig, layer_idx: int):
|
||||
super().__init__()
|
||||
self.hidden_size = config.hidden_size
|
||||
self.self_attn = Kanana2TinyAttention(config=config, layer_idx=layer_idx)
|
||||
self.mlp = Kanana2TinyMLP(config)
|
||||
self.input_layernorm = Kanana2TinyRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
self.post_attention_layernorm = Kanana2TinyRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
self.attention_type = config.layer_types[layer_idx]
|
||||
|
||||
@deprecate_kwarg("past_key_value", new_name="past_key_values", version="4.58")
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
position_embeddings_full: tuple[torch.Tensor, torch.Tensor],
|
||||
position_embeddings_sliding: tuple[torch.Tensor, torch.Tensor],
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
past_key_values: Optional[Cache] = None,
|
||||
use_cache: Optional[bool] = False,
|
||||
cache_position: Optional[torch.LongTensor] = None,
|
||||
**kwargs: Unpack[TransformersKwargs],
|
||||
) -> torch.Tensor:
|
||||
# Pick the right RoPE for this layer type.
|
||||
if self.attention_type == "sliding_attention":
|
||||
position_embeddings = position_embeddings_sliding
|
||||
else:
|
||||
position_embeddings = position_embeddings_full
|
||||
|
||||
residual = hidden_states
|
||||
hidden_states = self.input_layernorm(hidden_states)
|
||||
hidden_states, _ = self.self_attn(
|
||||
hidden_states=hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
position_ids=position_ids,
|
||||
past_key_values=past_key_values,
|
||||
use_cache=use_cache,
|
||||
cache_position=cache_position,
|
||||
position_embeddings=position_embeddings,
|
||||
**kwargs,
|
||||
)
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
residual = hidden_states
|
||||
hidden_states = self.post_attention_layernorm(hidden_states)
|
||||
hidden_states = self.mlp(hidden_states)
|
||||
hidden_states = residual + hidden_states
|
||||
return hidden_states
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Rotary embedding (driven by `config.rope_scaling` for the chosen attn type)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class Kanana2TinyRotaryEmbedding(nn.Module):
|
||||
"""Standard Qwen3-style rotary embedding.
|
||||
|
||||
The per-attention-type difference is encoded by the *config* passed in:
|
||||
callers construct two of these from views built via
|
||||
`_make_attention_specific_config` below.
|
||||
"""
|
||||
|
||||
inv_freq: torch.Tensor
|
||||
|
||||
def __init__(self, config: Kanana2TinyConfig, device=None):
|
||||
super().__init__()
|
||||
# BC: "rope_type" was originally "type"
|
||||
if hasattr(config, "rope_scaling") and isinstance(config.rope_scaling, dict):
|
||||
self.rope_type = config.rope_scaling.get("rope_type", config.rope_scaling.get("type", "default"))
|
||||
else:
|
||||
self.rope_type = "default"
|
||||
self.max_seq_len_cached = config.max_position_embeddings
|
||||
self.original_max_seq_len = config.max_position_embeddings
|
||||
|
||||
self.config = config
|
||||
# Resolve rope init across stock 5.4 (had 'default'), stock 5.5+
|
||||
# (renamed to 'proportional'), and Kakao-patched 5.3 (has neither;
|
||||
# falls through to our local unscaled-RoPE impl).
|
||||
self.rope_init_fn = _resolve_rope_init(self.rope_type)
|
||||
|
||||
inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device)
|
||||
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
||||
self.original_inv_freq = self.inv_freq
|
||||
|
||||
@staticmethod
|
||||
def compute_default_rope_parameters(config, device=None, seq_len=None):
|
||||
"""Stock transformers ≥5.9's ``modeling_utils._init_weights`` calls
|
||||
``module.compute_default_rope_parameters`` directly when ``rope_type
|
||||
== "default"`` (instead of looking it up in ``ROPE_INIT_FUNCTIONS``).
|
||||
This staticmethod has to exist on the class for that init pass to
|
||||
find it; the body is the same unscaled inv_freq computation we use
|
||||
as a fallback elsewhere.
|
||||
"""
|
||||
return _compute_default_rope_inv_freq(config, device=device, seq_len=seq_len)
|
||||
|
||||
@torch.no_grad()
|
||||
@dynamic_rope_update
|
||||
def forward(self, x, position_ids):
|
||||
inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1).to(x.device)
|
||||
position_ids_expanded = position_ids[:, None, :].float()
|
||||
|
||||
device_type = x.device.type if isinstance(x.device.type, str) and x.device.type != "mps" else "cpu"
|
||||
with torch.autocast(device_type=device_type, enabled=False):
|
||||
freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2)
|
||||
emb = torch.cat((freqs, freqs), dim=-1)
|
||||
cos = emb.cos() * self.attention_scaling
|
||||
sin = emb.sin() * self.attention_scaling
|
||||
|
||||
return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
|
||||
|
||||
|
||||
def _make_attention_specific_config(config: Kanana2TinyConfig, attention_type: str):
|
||||
"""Return a deep copy of `config` configured for a single attention type's
|
||||
RoPE. 4.57.1's `ROPE_INIT_FUNCTIONS` read `config.rope_theta` (top-level)
|
||||
and `config.rope_scaling` (a flat dict with `rope_type`/`factor`/...), so
|
||||
we flatten `config.rope_parameters[attention_type]` into that shape: pop
|
||||
`rope_theta` up to the top level, and leave the remaining keys in
|
||||
`rope_scaling`. For `rope_type='default'` this leaves a 1-key
|
||||
`{"rope_type": "default"}` dict, which `_validate_default_rope_parameters`
|
||||
accepts cleanly.
|
||||
"""
|
||||
if attention_type not in config.rope_parameters:
|
||||
raise KeyError(
|
||||
f"rope_parameters is missing entry for attention_type={attention_type!r}; "
|
||||
f"available keys: {list(config.rope_parameters.keys())}"
|
||||
)
|
||||
params = dict(config.rope_parameters[attention_type])
|
||||
new_config = copy.deepcopy(config)
|
||||
new_config.rope_theta = params.pop("rope_theta", config.rope_theta)
|
||||
new_config.rope_scaling = params
|
||||
return new_config
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pretrained model classes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@auto_docstring
|
||||
class Kanana2TinyPreTrainedModel(PreTrainedModel):
|
||||
config: Kanana2TinyConfig
|
||||
base_model_prefix = "model"
|
||||
supports_gradient_checkpointing = True
|
||||
_no_split_modules = ["Kanana2TinyDecoderLayer"]
|
||||
_skip_keys_device_placement = ["past_key_values"]
|
||||
_supports_flash_attn = True
|
||||
_supports_sdpa = True
|
||||
_supports_flex_attn = True
|
||||
|
||||
_can_compile_fullgraph = True
|
||||
_supports_attention_backend = True
|
||||
_can_record_outputs = {
|
||||
"hidden_states": Kanana2TinyDecoderLayer,
|
||||
"attentions": Kanana2TinyAttention,
|
||||
}
|
||||
|
||||
|
||||
@auto_docstring
|
||||
class Kanana2TinyModel(Kanana2TinyPreTrainedModel):
|
||||
def __init__(self, config: Kanana2TinyConfig):
|
||||
super().__init__(config)
|
||||
self.padding_idx = config.pad_token_id
|
||||
self.vocab_size = config.vocab_size
|
||||
|
||||
self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
|
||||
self.layers = nn.ModuleList(
|
||||
[Kanana2TinyDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
|
||||
)
|
||||
self.norm = Kanana2TinyRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
|
||||
# Two rotary embeddings, one per attention type. See the Gemma3
|
||||
# implementation for the same pattern.
|
||||
full_cfg = _make_attention_specific_config(config, "full_attention")
|
||||
self.rotary_emb_full = Kanana2TinyRotaryEmbedding(config=full_cfg)
|
||||
|
||||
if "sliding_attention" in config.layer_types:
|
||||
sliding_cfg = _make_attention_specific_config(config, "sliding_attention")
|
||||
self.rotary_emb_sliding = Kanana2TinyRotaryEmbedding(config=sliding_cfg)
|
||||
else:
|
||||
self.rotary_emb_sliding = None
|
||||
|
||||
# Backward-compat alias so any helper that expects `model.rotary_emb`
|
||||
# (e.g. some training-time monkey patches) still finds something.
|
||||
self.rotary_emb = self.rotary_emb_full
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
self.has_sliding_layers = "sliding_attention" in config.layer_types
|
||||
|
||||
self.post_init()
|
||||
|
||||
@check_model_inputs
|
||||
@auto_docstring
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.LongTensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
past_key_values: Optional[Cache] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
use_cache: Optional[bool] = None,
|
||||
cache_position: Optional[torch.LongTensor] = None,
|
||||
**kwargs: Unpack[TransformersKwargs],
|
||||
) -> BaseModelOutputWithPast:
|
||||
r"""
|
||||
cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*):
|
||||
Indices depicting the position of the input sequence tokens in the sequence. Used to
|
||||
update the cache in the correct position and to infer the complete sequence length.
|
||||
"""
|
||||
if (input_ids is None) ^ (inputs_embeds is not None):
|
||||
raise ValueError("You must specify exactly one of input_ids or inputs_embeds")
|
||||
|
||||
if inputs_embeds is None:
|
||||
inputs_embeds = self.embed_tokens(input_ids)
|
||||
|
||||
if use_cache and past_key_values is None:
|
||||
past_key_values = DynamicCache(config=self.config)
|
||||
|
||||
if cache_position is None:
|
||||
past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
|
||||
cache_position = torch.arange(
|
||||
past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device
|
||||
)
|
||||
|
||||
if position_ids is None:
|
||||
position_ids = cache_position.unsqueeze(0)
|
||||
|
||||
if not isinstance(causal_mask_mapping := attention_mask, dict):
|
||||
mask_kwargs = _filter_mask_kwargs({
|
||||
"config": self.config,
|
||||
_MASK_EMBEDS_KW: inputs_embeds,
|
||||
"attention_mask": attention_mask,
|
||||
"cache_position": cache_position,
|
||||
"past_key_values": past_key_values,
|
||||
"position_ids": position_ids,
|
||||
})
|
||||
causal_mask_mapping = {
|
||||
"full_attention": create_causal_mask(**mask_kwargs),
|
||||
}
|
||||
if self.has_sliding_layers:
|
||||
causal_mask_mapping["sliding_attention"] = create_sliding_window_causal_mask(**mask_kwargs)
|
||||
|
||||
hidden_states = inputs_embeds
|
||||
|
||||
position_embeddings_full = self.rotary_emb_full(hidden_states, position_ids)
|
||||
if self.rotary_emb_sliding is not None:
|
||||
position_embeddings_sliding = self.rotary_emb_sliding(hidden_states, position_ids)
|
||||
else:
|
||||
position_embeddings_sliding = position_embeddings_full
|
||||
|
||||
for decoder_layer in self.layers[: self.config.num_hidden_layers]:
|
||||
hidden_states = decoder_layer(
|
||||
hidden_states,
|
||||
position_embeddings_full=position_embeddings_full,
|
||||
position_embeddings_sliding=position_embeddings_sliding,
|
||||
attention_mask=causal_mask_mapping[decoder_layer.attention_type],
|
||||
position_ids=position_ids,
|
||||
past_key_values=past_key_values,
|
||||
use_cache=use_cache,
|
||||
cache_position=cache_position,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
return BaseModelOutputWithPast(
|
||||
last_hidden_state=hidden_states,
|
||||
past_key_values=past_key_values if use_cache else None,
|
||||
)
|
||||
|
||||
|
||||
@auto_docstring
|
||||
class Kanana2TinyForCausalLM(Kanana2TinyPreTrainedModel, GenerationMixin):
|
||||
# transformers v5 changed this from list to dict (mapping tied-key -> source-key).
|
||||
# The list form still works on v4. Use the dict form for forward-compatibility.
|
||||
_tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"}
|
||||
_tp_plan = {"lm_head": "colwise_rep"}
|
||||
_pp_plan = {"lm_head": (["hidden_states"], ["logits"])}
|
||||
|
||||
def __init__(self, config):
|
||||
super().__init__(config)
|
||||
self.model = Kanana2TinyModel(config)
|
||||
self.vocab_size = config.vocab_size
|
||||
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
|
||||
self.post_init()
|
||||
|
||||
@can_return_tuple
|
||||
@auto_docstring
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.LongTensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
past_key_values: Optional[Cache] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
labels: Optional[torch.LongTensor] = None,
|
||||
use_cache: Optional[bool] = None,
|
||||
cache_position: Optional[torch.LongTensor] = None,
|
||||
logits_to_keep: Union[int, torch.Tensor] = 0,
|
||||
**kwargs: Unpack[TransformersKwargs],
|
||||
) -> CausalLMOutputWithPast:
|
||||
r"""
|
||||
cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*):
|
||||
Indices depicting the position of the input sequence tokens in the sequence. Used to
|
||||
update the cache in the correct position and to infer the complete sequence length.
|
||||
labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
|
||||
Labels for computing the masked language modeling loss. Indices should either be in
|
||||
`[0, ..., config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices
|
||||
set to `-100` are ignored (masked); the loss is only computed for the tokens with
|
||||
labels in `[0, ..., config.vocab_size]`.
|
||||
"""
|
||||
outputs: BaseModelOutputWithPast = self.model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
position_ids=position_ids,
|
||||
past_key_values=past_key_values,
|
||||
inputs_embeds=inputs_embeds,
|
||||
use_cache=use_cache,
|
||||
cache_position=cache_position,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
hidden_states = outputs.last_hidden_state
|
||||
slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep
|
||||
logits = self.lm_head(hidden_states[:, slice_indices, :])
|
||||
|
||||
loss = None
|
||||
if labels is not None:
|
||||
loss = self.loss_function(logits=logits, labels=labels, vocab_size=self.config.vocab_size, **kwargs)
|
||||
|
||||
return CausalLMOutputWithPast(
|
||||
loss=loss,
|
||||
logits=logits,
|
||||
past_key_values=outputs.past_key_values,
|
||||
# Kanana2TinyModel.forward doesn't accumulate per-layer hidden_states even when
|
||||
# output_hidden_states=True; fall back to a 1-tuple of last_hidden_state so consumers
|
||||
# that index `hidden_states[-1]` (e.g. trl AutoModelForCausalLMWithValueHead) don't crash.
|
||||
hidden_states=outputs.hidden_states if outputs.hidden_states is not None else (outputs.last_hidden_state,),
|
||||
attentions=outputs.attentions,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"Kanana2TinyConfig",
|
||||
"Kanana2TinyForCausalLM",
|
||||
"Kanana2TinyModel",
|
||||
"Kanana2TinyPreTrainedModel",
|
||||
]
|
||||
73
sglang/config.json
Normal file
73
sglang/config.json
Normal file
@@ -0,0 +1,73 @@
|
||||
{
|
||||
"architectures": [
|
||||
"Qwen3ForCausalLM"
|
||||
],
|
||||
"attention_bias": false,
|
||||
"attention_dropout": 0.0,
|
||||
"dtype": "bfloat16",
|
||||
"head_dim": 128,
|
||||
"hidden_act": "silu",
|
||||
"hidden_size": 1280,
|
||||
"initializer_range": 0.02,
|
||||
"intermediate_size": 5760,
|
||||
"layer_types": [
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention"
|
||||
],
|
||||
"max_position_embeddings": 32768,
|
||||
"max_window_layers": 32,
|
||||
"model_type": "qwen3",
|
||||
"num_attention_heads": 32,
|
||||
"num_hidden_layers": 32,
|
||||
"num_key_value_heads": 8,
|
||||
"rms_norm_eps": 1e-06,
|
||||
"rope_theta": null,
|
||||
"rope_parameters": {
|
||||
"full_attention": {
|
||||
"rope_type": "yarn",
|
||||
"rope_theta": 10000,
|
||||
"factor": 40.0,
|
||||
"original_max_position_embeddings": 4096
|
||||
},
|
||||
"sliding_attention": {
|
||||
"rope_type": "default",
|
||||
"rope_theta": 10000.0
|
||||
}
|
||||
},
|
||||
"sliding_window": 1024,
|
||||
"tie_word_embeddings": true,
|
||||
"transformers_version": "4.57.6",
|
||||
"use_cache": true,
|
||||
"use_sliding_window": true,
|
||||
"vocab_size": 128256
|
||||
}
|
||||
558
sglang/qwen3.py
Normal file
558
sglang/qwen3.py
Normal file
@@ -0,0 +1,558 @@
|
||||
# Adapted from qwen2.py
|
||||
import logging
|
||||
from functools import partial
|
||||
from typing import Any, Dict, Iterable, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from sglang.srt.distributed import (
|
||||
get_pp_group,
|
||||
get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size,
|
||||
)
|
||||
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
|
||||
from sglang.srt.layers.dp_attention import get_attention_tp_rank, get_attention_tp_size
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import QKVParallelLinear, RowParallelLinear
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||
from sglang.srt.layers.pooler import Pooler, PoolingType
|
||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
from sglang.srt.layers.radix_attention import RadixAttention
|
||||
from sglang.srt.layers.rotary_embedding import get_rope
|
||||
from sglang.srt.layers.utils import PPMissingLayer, get_layer_id
|
||||
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
||||
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.models.qwen2 import Qwen2MLP as Qwen3MLP
|
||||
from sglang.srt.models.qwen2 import Qwen2Model
|
||||
from sglang.srt.utils import add_prefix, is_cuda
|
||||
|
||||
Qwen3Config = None
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
_is_cuda = is_cuda()
|
||||
|
||||
# Aligned with HF's implementation, using sliding window inclusive with the last token
|
||||
# SGLang assumes exclusive
|
||||
def get_attention_sliding_window_size(config):
|
||||
if getattr(config, "sliding_window", None) is not None:
|
||||
return config.sliding_window - 1
|
||||
else:
|
||||
return None
|
||||
|
||||
class Qwen3Attention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
num_kv_heads: int,
|
||||
layer_id: int = 0,
|
||||
rope_theta: float = 1000000,
|
||||
rope_scaling: Optional[Dict[str, Any]] = None,
|
||||
head_dim: Optional[int] = None,
|
||||
max_position_embeddings: int = 32768,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
rms_norm_eps: float = None,
|
||||
config=None,
|
||||
use_rope: bool = True,
|
||||
attention_bias: bool = False,
|
||||
prefix: str = "",
|
||||
alt_stream: Optional[torch.cuda.Stream] = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
self.total_num_heads = num_heads
|
||||
attn_tp_rank = get_attention_tp_rank()
|
||||
attn_tp_size = get_attention_tp_size()
|
||||
|
||||
assert self.total_num_heads % attn_tp_size == 0
|
||||
self.num_heads = self.total_num_heads // attn_tp_size
|
||||
self.total_num_kv_heads = num_kv_heads
|
||||
if self.total_num_kv_heads >= attn_tp_size:
|
||||
# Number of KV heads is greater than TP size, so we partition
|
||||
# the KV heads across multiple tensor parallel GPUs.
|
||||
assert self.total_num_kv_heads % attn_tp_size == 0
|
||||
else:
|
||||
# Number of KV heads is less than TP size, so we replicate
|
||||
# the KV heads across multiple tensor parallel GPUs.
|
||||
assert attn_tp_size % self.total_num_kv_heads == 0
|
||||
self.num_kv_heads = max(1, self.total_num_kv_heads // attn_tp_size)
|
||||
self.head_dim = head_dim or hidden_size // self.total_num_heads
|
||||
self.q_size = self.num_heads * self.head_dim
|
||||
self.kv_size = self.num_kv_heads * self.head_dim
|
||||
self.scaling = self.head_dim**-0.5
|
||||
self.rope_theta = rope_theta
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
self.tp_rank = get_tensor_model_parallel_rank()
|
||||
|
||||
self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps)
|
||||
self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps)
|
||||
|
||||
self.qkv_proj = QKVParallelLinear(
|
||||
hidden_size,
|
||||
self.head_dim,
|
||||
self.total_num_heads,
|
||||
self.total_num_kv_heads,
|
||||
bias=attention_bias,
|
||||
quant_config=quant_config,
|
||||
tp_rank=attn_tp_rank,
|
||||
tp_size=attn_tp_size,
|
||||
prefix=add_prefix("qkv_proj", prefix),
|
||||
)
|
||||
self.o_proj = RowParallelLinear(
|
||||
self.total_num_heads * self.head_dim,
|
||||
hidden_size,
|
||||
bias=attention_bias,
|
||||
quant_config=quant_config,
|
||||
tp_rank=attn_tp_rank,
|
||||
tp_size=attn_tp_size,
|
||||
reduce_results=False,
|
||||
prefix=add_prefix("o_proj", prefix),
|
||||
)
|
||||
|
||||
self.use_rope = use_rope
|
||||
self.rotary_emb = get_rope(
|
||||
self.head_dim,
|
||||
rotary_dim=self.head_dim,
|
||||
max_position=max_position_embeddings,
|
||||
base=rope_theta,
|
||||
rope_scaling=rope_scaling,
|
||||
)
|
||||
self.is_sliding = config.layer_types[layer_id] == "sliding_attention"
|
||||
self.attn = RadixAttention(
|
||||
self.num_heads,
|
||||
self.head_dim,
|
||||
self.scaling,
|
||||
num_kv_heads=self.num_kv_heads,
|
||||
layer_id=layer_id,
|
||||
sliding_window_size=(
|
||||
get_attention_sliding_window_size(config) if self.is_sliding else None
|
||||
),
|
||||
prefix=add_prefix("attn", prefix),
|
||||
)
|
||||
self.alt_stream = alt_stream
|
||||
|
||||
def _apply_qk_norm(
|
||||
self, q: torch.Tensor, k: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# overlap qk norm
|
||||
if self.alt_stream is not None and get_is_capture_mode():
|
||||
current_stream = torch.cuda.current_stream()
|
||||
self.alt_stream.wait_stream(current_stream)
|
||||
q_by_head = q.reshape(-1, self.head_dim)
|
||||
q_by_head = self.q_norm(q_by_head)
|
||||
with torch.cuda.stream(self.alt_stream):
|
||||
k_by_head = k.reshape(-1, self.head_dim)
|
||||
k_by_head = self.k_norm(k_by_head)
|
||||
current_stream.wait_stream(self.alt_stream)
|
||||
else:
|
||||
q_by_head = q.reshape(-1, self.head_dim)
|
||||
q_by_head = self.q_norm(q_by_head)
|
||||
k_by_head = k.reshape(-1, self.head_dim)
|
||||
k_by_head = self.k_norm(k_by_head)
|
||||
q = q_by_head.view(q.shape)
|
||||
k = k_by_head.view(k.shape)
|
||||
return q, k
|
||||
|
||||
def forward(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
) -> torch.Tensor:
|
||||
qkv, _ = self.qkv_proj(hidden_states)
|
||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||
q, k = self._apply_qk_norm(q, k)
|
||||
if self.use_rope:
|
||||
q, k = self.rotary_emb(positions, q, k)
|
||||
attn_output = self.attn(q, k, v, forward_batch)
|
||||
output, _ = self.o_proj(attn_output)
|
||||
return output
|
||||
|
||||
|
||||
class Qwen3DecoderLayer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
config: Qwen3Config,
|
||||
layer_id: int = 0,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
alt_stream: Optional[torch.cuda.Stream] = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.layer_id = layer_id
|
||||
self.hidden_size = config.hidden_size
|
||||
rope_theta = getattr(config, "rope_theta", 1000000)
|
||||
|
||||
if hasattr(config, "rope_parameters") and config.rope_parameters is not None:
|
||||
rope_scaling = config.rope_parameters
|
||||
else:
|
||||
rope_scaling = getattr(config, "rope_scaling", None)
|
||||
|
||||
self.use_rope = True
|
||||
|
||||
# nested → flat 변환
|
||||
if rope_scaling is not None:
|
||||
first_value = next(iter(rope_scaling.values()), None)
|
||||
if isinstance(first_value, dict):
|
||||
layer_type = config.layer_types[layer_id]
|
||||
if layer_type in rope_scaling:
|
||||
layer_rope = rope_scaling[layer_type]
|
||||
rope_theta = layer_rope.get("rope_theta", rope_theta)
|
||||
rope_scaling = layer_rope
|
||||
else:
|
||||
self.use_rope = False
|
||||
rope_scaling = None
|
||||
max_position_embeddings = getattr(config, "max_position_embeddings", 32768)
|
||||
head_dim = getattr(config, "head_dim", None)
|
||||
self.self_attn = Qwen3Attention(
|
||||
hidden_size=self.hidden_size,
|
||||
num_heads=config.num_attention_heads,
|
||||
num_kv_heads=config.num_key_value_heads,
|
||||
layer_id=layer_id,
|
||||
rope_theta=rope_theta,
|
||||
rope_scaling=rope_scaling,
|
||||
use_rope=self.use_rope,
|
||||
head_dim=head_dim,
|
||||
max_position_embeddings=max_position_embeddings,
|
||||
quant_config=quant_config,
|
||||
rms_norm_eps=config.rms_norm_eps,
|
||||
attention_bias=config.attention_bias,
|
||||
config=config,
|
||||
prefix=add_prefix("self_attn", prefix),
|
||||
alt_stream=alt_stream,
|
||||
)
|
||||
self.mlp = Qwen3MLP(
|
||||
hidden_size=self.hidden_size,
|
||||
intermediate_size=config.intermediate_size,
|
||||
hidden_act=config.hidden_act,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("mlp", prefix),
|
||||
)
|
||||
self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
self.post_attention_layernorm = RMSNorm(
|
||||
config.hidden_size, eps=config.rms_norm_eps
|
||||
)
|
||||
|
||||
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||
layer_id=layer_id,
|
||||
num_layers=config.num_hidden_layers,
|
||||
is_layer_sparse=False,
|
||||
is_previous_layer_sparse=False,
|
||||
)
|
||||
self.layer_communicator = LayerCommunicator(
|
||||
layer_scatter_modes=self.layer_scatter_modes,
|
||||
input_layernorm=self.input_layernorm,
|
||||
post_attention_layernorm=self.post_attention_layernorm,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
residual: Optional[torch.Tensor],
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# Self Attention
|
||||
hidden_states, residual = self.layer_communicator.prepare_attn(
|
||||
hidden_states, residual, forward_batch
|
||||
)
|
||||
if hidden_states.shape[0] != 0:
|
||||
hidden_states = self.self_attn(
|
||||
positions=positions,
|
||||
hidden_states=hidden_states,
|
||||
forward_batch=forward_batch,
|
||||
)
|
||||
|
||||
# Fully Connected
|
||||
hidden_states, residual = self.layer_communicator.prepare_mlp(
|
||||
hidden_states, residual, forward_batch
|
||||
)
|
||||
hidden_states = self.mlp(hidden_states)
|
||||
hidden_states, residual = self.layer_communicator.postprocess_layer(
|
||||
hidden_states, residual, forward_batch
|
||||
)
|
||||
return hidden_states, residual
|
||||
|
||||
|
||||
class Qwen3Model(Qwen2Model):
|
||||
def __init__(
|
||||
self,
|
||||
config: Qwen3Config,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
alt_stream = torch.cuda.Stream() if _is_cuda else None
|
||||
super().__init__(
|
||||
config=config,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix,
|
||||
decoder_layer_type=Qwen3DecoderLayer,
|
||||
alt_stream=alt_stream,
|
||||
)
|
||||
|
||||
|
||||
class Qwen3ForCausalLM(nn.Module):
|
||||
# BitandBytes specific attributes
|
||||
default_bitsandbytes_target_modules = [
|
||||
".gate_proj.",
|
||||
".down_proj.",
|
||||
".up_proj.",
|
||||
".q_proj.",
|
||||
".k_proj.",
|
||||
".v_proj.",
|
||||
".o_proj.",
|
||||
]
|
||||
bitsandbytes_stacked_params_mapping = {
|
||||
# shard_name, weight_name, index
|
||||
"q_proj": ("qkv_proj", 0),
|
||||
"k_proj": ("qkv_proj", 1),
|
||||
"v_proj": ("qkv_proj", 2),
|
||||
"gate_proj": ("gate_up_proj", 0),
|
||||
"up_proj": ("gate_up_proj", 1),
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Qwen3Config,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.pp_group = get_pp_group()
|
||||
self.config = config
|
||||
self.quant_config = quant_config
|
||||
self.model = Qwen3Model(
|
||||
config, quant_config=quant_config, prefix=add_prefix("model", prefix)
|
||||
)
|
||||
|
||||
# handle the lm head on different pp ranks
|
||||
if self.pp_group.is_last_rank:
|
||||
if self.pp_group.world_size == 1 and config.tie_word_embeddings:
|
||||
self.lm_head = self.model.embed_tokens
|
||||
else:
|
||||
self.lm_head = ParallelLMHead(
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
)
|
||||
else:
|
||||
# ranks other than the last rank will have a placeholder layer
|
||||
self.lm_head = PPMissingLayer()
|
||||
|
||||
# perform weight tying for PP
|
||||
if self.pp_group.world_size > 1 and config.tie_word_embeddings:
|
||||
if self.pp_group.is_first_rank:
|
||||
self.pp_group.send(
|
||||
self.model.embed_tokens.weight, dst=self.pp_group.last_rank
|
||||
)
|
||||
else:
|
||||
emb_token_weight = self.pp_group.recv(
|
||||
size=(config.vocab_size, config.hidden_size),
|
||||
dtype=next(self.model.parameters()).dtype,
|
||||
src=self.pp_group.first_rank,
|
||||
)
|
||||
self.lm_head.weight.copy_(emb_token_weight)
|
||||
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
|
||||
|
||||
# For EAGLE3 support
|
||||
self.capture_aux_hidden_states = False
|
||||
|
||||
def get_input_embeddings(self) -> nn.Embedding:
|
||||
return self.model.get_input_embeddings()
|
||||
|
||||
def get_attention_sliding_window_size(self):
|
||||
return get_attention_sliding_window_size(self.config)
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
input_embeds: torch.Tensor = None,
|
||||
get_embedding: bool = False,
|
||||
pp_proxy_tensors: Optional[PPProxyTensors] = None,
|
||||
) -> torch.Tensor:
|
||||
hidden_states = self.model(
|
||||
input_ids,
|
||||
positions,
|
||||
forward_batch,
|
||||
input_embeds,
|
||||
pp_proxy_tensors=pp_proxy_tensors,
|
||||
)
|
||||
|
||||
aux_hidden_states = None
|
||||
if self.capture_aux_hidden_states:
|
||||
hidden_states, aux_hidden_states = hidden_states
|
||||
|
||||
if self.pp_group.is_last_rank:
|
||||
if not get_embedding:
|
||||
return self.logits_processor(
|
||||
input_ids,
|
||||
hidden_states,
|
||||
self.lm_head,
|
||||
forward_batch,
|
||||
aux_hidden_states,
|
||||
)
|
||||
else:
|
||||
return self.pooler(hidden_states, forward_batch)
|
||||
else:
|
||||
return hidden_states
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_split_prefill(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
split_interval: Tuple[int, int], # [start, end) 0-based
|
||||
input_embeds: torch.Tensor = None,
|
||||
):
|
||||
start, end = split_interval
|
||||
# embed
|
||||
if start == 0:
|
||||
if input_embeds is None:
|
||||
forward_batch.hidden_states = self.model.embed_tokens(input_ids)
|
||||
else:
|
||||
forward_batch.hidden_states = input_embeds
|
||||
# decoder layer
|
||||
for i in range(start, end):
|
||||
layer = self.model.layers[i]
|
||||
forward_batch.hidden_states, forward_batch.residual = layer(
|
||||
positions,
|
||||
forward_batch.hidden_states,
|
||||
forward_batch,
|
||||
forward_batch.residual,
|
||||
)
|
||||
|
||||
if end == self.model.config.num_hidden_layers:
|
||||
# norm
|
||||
hidden_states, _ = self.model.norm(
|
||||
forward_batch.hidden_states, forward_batch.residual
|
||||
)
|
||||
forward_batch.hidden_states = hidden_states
|
||||
# logits process
|
||||
result = self.logits_processor(
|
||||
input_ids, forward_batch.hidden_states, self.lm_head, forward_batch
|
||||
)
|
||||
else:
|
||||
result = None
|
||||
|
||||
return result
|
||||
|
||||
@property
|
||||
def start_layer(self):
|
||||
return self.model.start_layer
|
||||
|
||||
@property
|
||||
def end_layer(self):
|
||||
return self.model.end_layer
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
("qkv_proj", "q_proj", "q"),
|
||||
("qkv_proj", "k_proj", "k"),
|
||||
("qkv_proj", "v_proj", "v"),
|
||||
("gate_up_proj", "gate_proj", 0),
|
||||
("gate_up_proj", "up_proj", 1),
|
||||
]
|
||||
|
||||
params_dict = dict(self.named_parameters())
|
||||
for name, loaded_weight in weights:
|
||||
if "Embedding" in self.config.name_or_path:
|
||||
name = add_prefix(name, "model")
|
||||
layer_id = get_layer_id(name)
|
||||
if (
|
||||
layer_id is not None
|
||||
and hasattr(self.model, "start_layer")
|
||||
and (
|
||||
layer_id < self.model.start_layer
|
||||
or layer_id >= self.model.end_layer
|
||||
)
|
||||
):
|
||||
continue
|
||||
|
||||
if "rotary_emb.inv_freq" in name or "projector" in name:
|
||||
continue
|
||||
if "rotary_emb.cos_cached" in name or "rotary_emb.sin_cached" in name:
|
||||
# Models trained using ColossalAI may include these tensors in
|
||||
# the checkpoint. Skip them.
|
||||
continue
|
||||
if self.config.tie_word_embeddings and "lm_head.weight" in name:
|
||||
if self.pp_group.world_size > 1 and self.pp_group.is_last_rank:
|
||||
# Handle pp weight tying here
|
||||
# find the embed_tokens.weight in the weights
|
||||
embed_token_weights = next(
|
||||
filter(lambda x: x[0] == "model.embed_tokens.weight", weights)
|
||||
)[1]
|
||||
loaded_weight = embed_token_weights
|
||||
else:
|
||||
continue
|
||||
if name.startswith("model.vision_tower") and name not in params_dict:
|
||||
continue
|
||||
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
param = params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
break
|
||||
else:
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
|
||||
if name in params_dict.keys():
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(
|
||||
param, "weight_loader", default_weight_loader
|
||||
)
|
||||
weight_loader(param, loaded_weight)
|
||||
else:
|
||||
logger.warning(f"Parameter {name} not found in params_dict")
|
||||
|
||||
def get_embed_and_head(self):
|
||||
return self.model.embed_tokens.weight, self.lm_head.weight
|
||||
|
||||
def set_embed_and_head(self, embed, head):
|
||||
del self.model.embed_tokens.weight
|
||||
del self.lm_head.weight
|
||||
self.model.embed_tokens.weight = embed
|
||||
self.lm_head.weight = head
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
def load_kv_cache_scales(self, quantization_param_path: str) -> None:
|
||||
self.model.load_kv_cache_scales(quantization_param_path)
|
||||
|
||||
def set_eagle3_layers_to_capture(self, layer_ids: Optional[List[int]] = None):
|
||||
if not self.pp_group.is_last_rank:
|
||||
return
|
||||
|
||||
self.capture_aux_hidden_states = True
|
||||
if layer_ids is None:
|
||||
num_layers = self.config.num_hidden_layers
|
||||
self.model.layers_to_capture = [
|
||||
2,
|
||||
num_layers // 2,
|
||||
num_layers - 3,
|
||||
] # Specific layers for EAGLE3 support
|
||||
else:
|
||||
self.model.layers_to_capture = [val + 1 for val in layer_ids]
|
||||
|
||||
|
||||
EntryClass = Qwen3ForCausalLM
|
||||
641401
tokenizer.json
Normal file
641401
tokenizer.json
Normal file
File diff suppressed because it is too large
Load Diff
2063
tokenizer_config.json
Normal file
2063
tokenizer_config.json
Normal file
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user