Continual finetuning of microsoft/elem2design on our small datasets --> Upload full checkpoint including LoRA and projectors
Browse files- README.md +70 -103
- adapter_config.json +4 -8
- config.json +51 -0
- mm_projector.bin +3 -0
- rng_state_0.pth +3 -0
- rng_state_1.pth +3 -0
- rng_state_2.pth +3 -0
- rng_state_3.pth +3 -0
- rng_state_4.pth +3 -0
- rng_state_5.pth +3 -0
- rng_state_6.pth +3 -0
- rng_state_7.pth +3 -0
- scheduler.pt +3 -0
- tokenizer.json +2 -2
- tokenizer_config.json +1 -10
- trainer_state.json +0 -0
- training_args.bin +3 -0
- zero_to_fp32.py +587 -0
README.md
CHANGED
|
@@ -1,37 +1,27 @@
|
|
| 1 |
---
|
| 2 |
-
|
| 3 |
-
tags: []
|
| 4 |
---
|
| 5 |
|
| 6 |
-
# Model Card for Model ID
|
| 7 |
-
|
| 8 |
-
<!-- Provide a quick summary of what the model is/does. -->
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
| 12 |
## Model Details
|
| 13 |
|
| 14 |
### Model Description
|
| 15 |
|
| 16 |
-
|
|
|
|
| 17 |
|
| 18 |
-
This is the model card of a 🤗 transformers model that has been pushed on the Hub. This model card has been automatically generated.
|
| 19 |
|
| 20 |
-
- **Developed by:**
|
| 21 |
-
- **
|
| 22 |
-
- **
|
| 23 |
-
- **
|
| 24 |
-
- **
|
| 25 |
-
- **License:** [More Information Needed]
|
| 26 |
-
- **Finetuned from model [optional]:** [More Information Needed]
|
| 27 |
|
| 28 |
-
### Model Sources
|
| 29 |
|
| 30 |
<!-- Provide the basic links for the model. -->
|
| 31 |
|
| 32 |
-
- **Repository:**
|
| 33 |
-
- **Paper
|
| 34 |
-
- **Demo [optional]:** [More Information Needed]
|
| 35 |
|
| 36 |
## Uses
|
| 37 |
|
|
@@ -41,37 +31,51 @@ This is the model card of a 🤗 transformers model that has been pushed on the
|
|
| 41 |
|
| 42 |
<!-- This section is for the model use without fine-tuning or plugging into a larger ecosystem/app. -->
|
| 43 |
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
### Downstream Use [optional]
|
| 47 |
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
[More Information Needed]
|
| 51 |
|
| 52 |
### Out-of-Scope Use
|
| 53 |
|
| 54 |
<!-- This section addresses misuse, malicious use, and uses that the model will not work well for. -->
|
| 55 |
|
| 56 |
-
|
|
|
|
|
|
|
| 57 |
|
| 58 |
-
##
|
| 59 |
|
| 60 |
<!-- This section is meant to convey both technical and sociotechnical limitations. -->
|
| 61 |
|
| 62 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 63 |
|
| 64 |
### Recommendations
|
| 65 |
|
| 66 |
<!-- This section is meant to convey recommendations with respect to the bias, risk, and technical limitations. -->
|
| 67 |
|
| 68 |
-
|
| 69 |
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
Use the code below to get started with the model.
|
| 73 |
|
| 74 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 75 |
|
| 76 |
## Training Details
|
| 77 |
|
|
@@ -79,64 +83,62 @@ Use the code below to get started with the model.
|
|
| 79 |
|
| 80 |
<!-- This should link to a Dataset Card, perhaps with a short stub of information on what the training data is all about as well as documentation related to data pre-processing or additional filtering. -->
|
| 81 |
|
| 82 |
-
|
| 83 |
|
| 84 |
### Training Procedure
|
| 85 |
|
| 86 |
<!-- This relates heavily to the Technical Specifications. Content here should link to that section when it is relevant to the training procedure. -->
|
| 87 |
|
| 88 |
-
#### Preprocessing
|
| 89 |
|
| 90 |
-
|
| 91 |
|
| 92 |
|
| 93 |
#### Training Hyperparameters
|
| 94 |
|
| 95 |
-
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 96 |
|
| 97 |
-
#### Speeds, Sizes, Times
|
| 98 |
|
| 99 |
<!-- This section provides information about throughput, start/end time, checkpoint size if relevant, etc. -->
|
| 100 |
|
| 101 |
-
|
|
|
|
|
|
|
| 102 |
|
| 103 |
## Evaluation
|
| 104 |
|
| 105 |
<!-- This section describes the evaluation protocols and provides the results. -->
|
| 106 |
|
| 107 |
-
### Testing Data, Factors & Metrics
|
| 108 |
|
| 109 |
#### Testing Data
|
| 110 |
|
| 111 |
<!-- This should link to a Dataset Card if possible. -->
|
| 112 |
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
#### Factors
|
| 116 |
-
|
| 117 |
-
<!-- These are the things the evaluation is disaggregating by, e.g., subpopulations or domains. -->
|
| 118 |
|
| 119 |
-
[More Information Needed]
|
| 120 |
|
| 121 |
#### Metrics
|
| 122 |
|
| 123 |
<!-- These are the evaluation metrics being used, ideally with a description of why. -->
|
| 124 |
|
| 125 |
-
[More Information Needed]
|
| 126 |
-
|
| 127 |
-
### Results
|
| 128 |
|
| 129 |
-
|
| 130 |
|
| 131 |
-
|
| 132 |
|
| 133 |
|
|
|
|
| 134 |
|
| 135 |
-
|
| 136 |
|
| 137 |
-
<!-- Relevant interpretability work for the model goes here -->
|
| 138 |
|
| 139 |
-
[More Information Needed]
|
| 140 |
|
| 141 |
## Environmental Impact
|
| 142 |
|
|
@@ -144,56 +146,21 @@ Use the code below to get started with the model.
|
|
| 144 |
|
| 145 |
Carbon emissions can be estimated using the [Machine Learning Impact calculator](https://mlco2.github.io/impact#compute) presented in [Lacoste et al. (2019)](https://arxiv.org/abs/1910.09700).
|
| 146 |
|
| 147 |
-
- **Hardware Type:** [More Information Needed]
|
| 148 |
-
- **Hours used:** [More Information Needed]
|
| 149 |
-
- **Cloud Provider:** [More Information Needed]
|
| 150 |
-
- **Compute Region:** [More Information Needed]
|
| 151 |
-
- **Carbon Emitted:** [More Information Needed]
|
| 152 |
-
|
| 153 |
-
## Technical Specifications [optional]
|
| 154 |
-
|
| 155 |
-
### Model Architecture and Objective
|
| 156 |
-
|
| 157 |
-
[More Information Needed]
|
| 158 |
-
|
| 159 |
-
### Compute Infrastructure
|
| 160 |
-
|
| 161 |
-
[More Information Needed]
|
| 162 |
|
| 163 |
-
##
|
| 164 |
-
|
| 165 |
-
[More Information Needed]
|
| 166 |
-
|
| 167 |
-
#### Software
|
| 168 |
-
|
| 169 |
-
[More Information Needed]
|
| 170 |
-
|
| 171 |
-
## Citation [optional]
|
| 172 |
|
| 173 |
<!-- If there is a paper or blog post introducing the model, the APA and Bibtex information for that should go in this section. -->
|
| 174 |
-
|
| 175 |
-
|
| 176 |
-
|
| 177 |
-
|
| 178 |
-
|
| 179 |
-
|
| 180 |
-
|
| 181 |
-
|
| 182 |
-
|
| 183 |
-
## Glossary [optional]
|
| 184 |
-
|
| 185 |
-
<!-- If relevant, include terms and calculations in this section that can help readers understand the model or model card. -->
|
| 186 |
-
|
| 187 |
-
[More Information Needed]
|
| 188 |
-
|
| 189 |
-
## More Information [optional]
|
| 190 |
-
|
| 191 |
-
[More Information Needed]
|
| 192 |
-
|
| 193 |
-
## Model Card Authors [optional]
|
| 194 |
-
|
| 195 |
-
[More Information Needed]
|
| 196 |
|
| 197 |
## Model Card Contact
|
| 198 |
|
| 199 |
-
|
|
|
|
|
|
|
|
|
| 1 |
---
|
| 2 |
+
license: mit
|
|
|
|
| 3 |
---
|
| 4 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 5 |
## Model Details
|
| 6 |
|
| 7 |
### Model Description
|
| 8 |
|
| 9 |
+
This model aims to compose user-provided graphic elements into a pleasing graphical design. It takes graphic elements (i.e., the images and texts) from users as input and generates the position, color and font information of each element as output.
|
| 10 |
+
|
| 11 |
|
|
|
|
| 12 |
|
| 13 |
+
- **Developed by:** Jiawei Lin, Shizhao Sun, Danqing Huang, Ting Liu, Ji Li and Jiang Bian
|
| 14 |
+
- **Model type:** Large Language Models
|
| 15 |
+
- **Language(s):** Python
|
| 16 |
+
- **License:** MIT
|
| 17 |
+
- **Finetuned from model:** Llama-3.1-8B
|
|
|
|
|
|
|
| 18 |
|
| 19 |
+
### Model Sources
|
| 20 |
|
| 21 |
<!-- Provide the basic links for the model. -->
|
| 22 |
|
| 23 |
+
- **Repository:** https://github.com/microsoft/elem2design
|
| 24 |
+
- **Paper:** https://arxiv.org/abs/2412.19712
|
|
|
|
| 25 |
|
| 26 |
## Uses
|
| 27 |
|
|
|
|
| 31 |
|
| 32 |
<!-- This section is for the model use without fine-tuning or plugging into a larger ecosystem/app. -->
|
| 33 |
|
| 34 |
+
Compose user-provided graphic elements (i.e., images and texts) into a pleasing graphic design.
|
|
|
|
|
|
|
| 35 |
|
| 36 |
+
Elem2Design is being shared with the research community to facilitate reproduction of our results and foster further research in this area.
|
|
|
|
|
|
|
| 37 |
|
| 38 |
### Out-of-Scope Use
|
| 39 |
|
| 40 |
<!-- This section addresses misuse, malicious use, and uses that the model will not work well for. -->
|
| 41 |
|
| 42 |
+
We do not recommend using Elem2Design in commercial or real-world applications without further testing and development. It is being released for research purposes.
|
| 43 |
+
|
| 44 |
+
Use in any manner that violates applicable laws or regulations.
|
| 45 |
|
| 46 |
+
## Risks and Limitations
|
| 47 |
|
| 48 |
<!-- This section is meant to convey both technical and sociotechnical limitations. -->
|
| 49 |
|
| 50 |
+
Elem2Design inherits any biases, errors, or omissions produced by its base model. Developers are advised to choose an appropriate base LLM/MLLM carefully, depending on the intended use case.
|
| 51 |
+
|
| 52 |
+
Elem2Design uses the Llama model. See https://huggingface.co/meta-llama/Llama-3.1-8B to understand the capabilities and limitations of this model.
|
| 53 |
+
|
| 54 |
+
As the model is fine-tuned on very specific data about design composition, it is unlikely to generate information other than position, color and font. However, this is possible. It is more likely to happen when instructions unrelated to graphic design composition, e.g., how has the social media influenced our daily life, are fed into the model.
|
| 55 |
+
|
| 56 |
+
Graphic designs generated by Elem2Design may not be technically accurate or meet user specifications in all cases. Users are responsible for assessing the acceptability of generated content for each intended use case.
|
| 57 |
+
|
| 58 |
+
Elem2Design was developed for research and experimental purposes. Further testing and validation are needed before considering its application in commercial or real-world scenarios.
|
| 59 |
|
| 60 |
### Recommendations
|
| 61 |
|
| 62 |
<!-- This section is meant to convey recommendations with respect to the bias, risk, and technical limitations. -->
|
| 63 |
|
| 64 |
+
Please only provide the images and texts that you want to show on the graphic design to the model.
|
| 65 |
|
| 66 |
+
Users are responsible for sourcing their content legally and ethically. This could include securing appropriate copy rights, ensuring consent for use of images of people, and/or the anonymization of data prior to use in research.
|
|
|
|
|
|
|
| 67 |
|
| 68 |
+
## How to Get Started with the Model
|
| 69 |
+
```
|
| 70 |
+
python llava/infer/infer.py \
|
| 71 |
+
--model_name_or_path /path/to/model/checkpoint-xxxx \
|
| 72 |
+
--data_path /path/to/data/test.json \
|
| 73 |
+
--image_folder /path/to/crello_images \
|
| 74 |
+
--output_dir /path/to/output_dir \
|
| 75 |
+
--start_layer_index 0 \
|
| 76 |
+
--end_layer_index 4
|
| 77 |
+
```
|
| 78 |
+
For more information, please visit our GitHub repo: https://github.com/microsoft/elem2design.
|
| 79 |
|
| 80 |
## Training Details
|
| 81 |
|
|
|
|
| 83 |
|
| 84 |
<!-- This should link to a Dataset Card, perhaps with a short stub of information on what the training data is all about as well as documentation related to data pre-processing or additional filtering. -->
|
| 85 |
|
| 86 |
+
The training data is from an open-source dataset (https://huggingface.co/datasets/cyberagent/crello).
|
| 87 |
|
| 88 |
### Training Procedure
|
| 89 |
|
| 90 |
<!-- This relates heavily to the Technical Specifications. Content here should link to that section when it is relevant to the training procedure. -->
|
| 91 |
|
| 92 |
+
#### Preprocessing
|
| 93 |
|
| 94 |
+
The training samples with more than 25 design elements are filtered out to maintain a limited sequence length and thereby improve training efficiency.
|
| 95 |
|
| 96 |
|
| 97 |
#### Training Hyperparameters
|
| 98 |
|
| 99 |
+
- Learning rate: 2e-4
|
| 100 |
+
|
| 101 |
+
- Global batch size: 128
|
| 102 |
+
|
| 103 |
+
- Number of training steps: 7000
|
| 104 |
+
|
| 105 |
+
- Rank and alpha of LoRA: 32 and 64
|
| 106 |
|
| 107 |
+
#### Speeds, Sizes, Times
|
| 108 |
|
| 109 |
<!-- This section provides information about throughput, start/end time, checkpoint size if relevant, etc. -->
|
| 110 |
|
| 111 |
+
- Llama-3.1-8B: 8B parameters
|
| 112 |
+
|
| 113 |
+
- CLIP ViT-Large-Patch14: 428M parameters
|
| 114 |
|
| 115 |
## Evaluation
|
| 116 |
|
| 117 |
<!-- This section describes the evaluation protocols and provides the results. -->
|
| 118 |
|
|
|
|
| 119 |
|
| 120 |
#### Testing Data
|
| 121 |
|
| 122 |
<!-- This should link to a Dataset Card if possible. -->
|
| 123 |
|
| 124 |
+
The testing data is from an open-source dataset (https://huggingface.co/datasets/cyberagent/crello).
|
|
|
|
|
|
|
|
|
|
|
|
|
| 125 |
|
|
|
|
| 126 |
|
| 127 |
#### Metrics
|
| 128 |
|
| 129 |
<!-- These are the evaluation metrics being used, ideally with a description of why. -->
|
| 130 |
|
|
|
|
|
|
|
|
|
|
| 131 |
|
| 132 |
+
- Overall metrics. We use a robust proxy model (https://huggingface.co/llava-hf/llava-onevision-qwen2-7b-ov-hf) for comprehensive evaluation from five aspects: (i) design and layout, (ii) content relevance, (iii) typography and color, (iv) graphics and images, and (v) innovation and originality. We use the same prompts as presented in COLE [[1](https://arxiv.org/abs/2311.16974)].
|
| 133 |
|
| 134 |
+
- Geometry-related metrics. These metrics focus purely on the geometric attributes of elements without considering their content, including element validity (Val), Overlap (Ove), Alignment (Ali) and underlay effectiveness (Undl, Unds) [[2](https://arxiv.org/abs/2303.15937 )][[3](https://arxiv.org/pdf/2404.00995 )].
|
| 135 |
|
| 136 |
|
| 137 |
+
### Results
|
| 138 |
|
| 139 |
+
We use prior work FlexDM [[1](https://arxiv.org/pdf/2303.18248)] and prompting GPT-4o [[2](https://platform.openai.com/docs/models#gpt-4o)] as baselines. In comparison, Elem2Design demonstrates superior performance across nearly all metrics. For example, on overall metrics, Elem2Design achieves 8.08, 7.92, 8.00, 7.82 and 6.98 on the evaluated five aspects. For another example, regarding geometry-related metrics, Elem2Design and FlexDM achieve Ove score of 0.0865 and 0.3242 respectively, indicting that Elem2Design effectively addresses the overlap issue whereas FlexDM encounters difficulties in this area. See Table 1 for the complete evaluation in our paper (https://arxiv.org/pdf/2412.19712)
|
| 140 |
|
|
|
|
| 141 |
|
|
|
|
| 142 |
|
| 143 |
## Environmental Impact
|
| 144 |
|
|
|
|
| 146 |
|
| 147 |
Carbon emissions can be estimated using the [Machine Learning Impact calculator](https://mlco2.github.io/impact#compute) presented in [Lacoste et al. (2019)](https://arxiv.org/abs/1910.09700).
|
| 148 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 149 |
|
| 150 |
+
## Citation
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 151 |
|
| 152 |
<!-- If there is a paper or blog post introducing the model, the APA and Bibtex information for that should go in this section. -->
|
| 153 |
+
```
|
| 154 |
+
@InProceedings{lin2024elements,
|
| 155 |
+
title={From Elements to Design: A Layered Approach for Automatic Graphic Design Composition},
|
| 156 |
+
author={Lin, Jiawei and Sun, Shizhao and Huang, Danqing and Liu, Ting and Li, Ji and Bian, Jiang},
|
| 157 |
+
booktitle={CVPR},
|
| 158 |
+
year={2025}
|
| 159 |
+
}
|
| 160 |
+
```
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 161 |
|
| 162 |
## Model Card Contact
|
| 163 |
|
| 164 |
+
We welcome feedback and collaboration from our audience. If you have suggestions, questions, or observe unexpected/offensive behavior in our technology, please contact us at Shizhao Sun, shizsu@microsoft.com.
|
| 165 |
+
|
| 166 |
+
If the team receives reports of undesired behavior or identifies issues independently, we will update this repository with appropriate mitigations.
|
adapter_config.json
CHANGED
|
@@ -3,7 +3,6 @@
|
|
| 3 |
"auto_mapping": null,
|
| 4 |
"base_model_name_or_path": "meta-llama/Llama-3.1-8B",
|
| 5 |
"bias": "none",
|
| 6 |
-
"corda_config": null,
|
| 7 |
"eva_config": null,
|
| 8 |
"exclude_modules": null,
|
| 9 |
"fan_in_fan_out": false,
|
|
@@ -20,22 +19,19 @@
|
|
| 20 |
"megatron_core": "megatron.core",
|
| 21 |
"modules_to_save": null,
|
| 22 |
"peft_type": "LORA",
|
| 23 |
-
"qalora_group_size": 16,
|
| 24 |
"r": 32,
|
| 25 |
"rank_pattern": {},
|
| 26 |
"revision": null,
|
| 27 |
"target_modules": [
|
| 28 |
-
"q_proj",
|
| 29 |
-
"up_proj",
|
| 30 |
"o_proj",
|
|
|
|
| 31 |
"down_proj",
|
|
|
|
| 32 |
"v_proj",
|
| 33 |
-
"
|
| 34 |
-
"
|
| 35 |
],
|
| 36 |
"task_type": "CAUSAL_LM",
|
| 37 |
-
"trainable_token_indices": null,
|
| 38 |
"use_dora": false,
|
| 39 |
-
"use_qalora": false,
|
| 40 |
"use_rslora": false
|
| 41 |
}
|
|
|
|
| 3 |
"auto_mapping": null,
|
| 4 |
"base_model_name_or_path": "meta-llama/Llama-3.1-8B",
|
| 5 |
"bias": "none",
|
|
|
|
| 6 |
"eva_config": null,
|
| 7 |
"exclude_modules": null,
|
| 8 |
"fan_in_fan_out": false,
|
|
|
|
| 19 |
"megatron_core": "megatron.core",
|
| 20 |
"modules_to_save": null,
|
| 21 |
"peft_type": "LORA",
|
|
|
|
| 22 |
"r": 32,
|
| 23 |
"rank_pattern": {},
|
| 24 |
"revision": null,
|
| 25 |
"target_modules": [
|
|
|
|
|
|
|
| 26 |
"o_proj",
|
| 27 |
+
"gate_proj",
|
| 28 |
"down_proj",
|
| 29 |
+
"q_proj",
|
| 30 |
"v_proj",
|
| 31 |
+
"k_proj",
|
| 32 |
+
"up_proj"
|
| 33 |
],
|
| 34 |
"task_type": "CAUSAL_LM",
|
|
|
|
| 35 |
"use_dora": false,
|
|
|
|
| 36 |
"use_rslora": false
|
| 37 |
}
|
config.json
ADDED
|
@@ -0,0 +1,51 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_name_or_path": "meta-llama/Llama-3.1-8B",
|
| 3 |
+
"architectures": [
|
| 4 |
+
"LlamaForCausalLM"
|
| 5 |
+
],
|
| 6 |
+
"attention_bias": false,
|
| 7 |
+
"attention_dropout": 0.0,
|
| 8 |
+
"bos_token_id": 128000,
|
| 9 |
+
"eos_token_id": 128001,
|
| 10 |
+
"freeze_mm_mlp_adapter": false,
|
| 11 |
+
"hidden_act": "silu",
|
| 12 |
+
"hidden_size": 4096,
|
| 13 |
+
"image_aspect_ratio": "pad",
|
| 14 |
+
"initializer_range": 0.02,
|
| 15 |
+
"intermediate_size": 14336,
|
| 16 |
+
"max_position_embeddings": 131072,
|
| 17 |
+
"mlp_bias": false,
|
| 18 |
+
"mm_hidden_size": 1024,
|
| 19 |
+
"mm_patch_merge_type": "flat",
|
| 20 |
+
"mm_projector_lr": 0.0002,
|
| 21 |
+
"mm_projector_type": "mlp2x_gelu",
|
| 22 |
+
"mm_use_im_patch_token": false,
|
| 23 |
+
"mm_use_im_start_end": false,
|
| 24 |
+
"mm_vision_select_feature": "cls_pooling",
|
| 25 |
+
"mm_vision_select_layer": -2,
|
| 26 |
+
"mm_vision_tower": "openai/clip-vit-large-patch14-336",
|
| 27 |
+
"model_type": "llava_llama",
|
| 28 |
+
"num_attention_heads": 32,
|
| 29 |
+
"num_hidden_layers": 32,
|
| 30 |
+
"num_key_value_heads": 8,
|
| 31 |
+
"num_pooling_token": 4,
|
| 32 |
+
"pretraining_tp": 1,
|
| 33 |
+
"rms_norm_eps": 1e-05,
|
| 34 |
+
"rope_scaling": {
|
| 35 |
+
"factor": 8.0,
|
| 36 |
+
"high_freq_factor": 4.0,
|
| 37 |
+
"low_freq_factor": 1.0,
|
| 38 |
+
"original_max_position_embeddings": 8192,
|
| 39 |
+
"rope_type": "llama3"
|
| 40 |
+
},
|
| 41 |
+
"rope_theta": 500000.0,
|
| 42 |
+
"tie_word_embeddings": false,
|
| 43 |
+
"tokenizer_model_max_length": 15000,
|
| 44 |
+
"tokenizer_padding_side": "right",
|
| 45 |
+
"torch_dtype": "bfloat16",
|
| 46 |
+
"transformers_version": "4.44.2",
|
| 47 |
+
"tune_mm_mlp_adapter": false,
|
| 48 |
+
"use_cache": false,
|
| 49 |
+
"use_mm_proj": true,
|
| 50 |
+
"vocab_size": 128256
|
| 51 |
+
}
|
mm_projector.bin
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:f07fc46020e9555e7413e71875be60393c40b8b527424097dc0e29f7e1838ae7
|
| 3 |
+
size 41961592
|
rng_state_0.pth
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:d86d6316667c268891d1b90bf0fa259106e90311c4e9e71f74b41221b5f80174
|
| 3 |
+
size 15920
|
rng_state_1.pth
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:c9042a297c8cd78a2ef58e3a81fbd6cd68503930cde10c5852bf7051b8be55c3
|
| 3 |
+
size 15920
|
rng_state_2.pth
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:bfa4ce2a94217fd28e70efd6419c02692bae7d6dd4aa8482add5bd2555e1880e
|
| 3 |
+
size 15920
|
rng_state_3.pth
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:d7006071f420730e4c6ee0a77699b4d64d865651c4ee538c34bbcd04bb3fa145
|
| 3 |
+
size 15920
|
rng_state_4.pth
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:400a062ad2874d0bc502723cc830f2a33c07cd943311ebcaef6c82559e5665e3
|
| 3 |
+
size 15920
|
rng_state_5.pth
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:ea0008e4875c09ed7d2b898f65ceb0538fb99df3599aaf73ba0b8dd507516e12
|
| 3 |
+
size 15920
|
rng_state_6.pth
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:fb1d4dbd46a53af129b268aa4e8cf7bce675780d53e079d10825b5886883cda8
|
| 3 |
+
size 15920
|
rng_state_7.pth
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:1c17e34bb69a668b351c641d1d879a0d3dc77d273f62a48b30ca63d6f90b5109
|
| 3 |
+
size 15920
|
scheduler.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:bea2f2fdffa3af1d5eb2404824dc82cb147ff3a7520ebfd8bae0a9bd2021a690
|
| 3 |
+
size 1064
|
tokenizer.json
CHANGED
|
@@ -1,3 +1,3 @@
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
-
oid sha256:
|
| 3 |
-
size
|
|
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:79e3e522635f3171300913bb421464a87de6222182a0570b9b2ccba2a964b2b4
|
| 3 |
+
size 9085657
|
tokenizer_config.json
CHANGED
|
@@ -1,13 +1,5 @@
|
|
| 1 |
{
|
| 2 |
"added_tokens_decoder": {
|
| 3 |
-
"0": {
|
| 4 |
-
"content": "!",
|
| 5 |
-
"lstrip": false,
|
| 6 |
-
"normalized": false,
|
| 7 |
-
"rstrip": false,
|
| 8 |
-
"single_word": false,
|
| 9 |
-
"special": true
|
| 10 |
-
},
|
| 11 |
"128000": {
|
| 12 |
"content": "<|begin_of_text|>",
|
| 13 |
"lstrip": false,
|
|
@@ -2060,7 +2052,6 @@
|
|
| 2060 |
"bos_token": "<|begin_of_text|>",
|
| 2061 |
"clean_up_tokenization_spaces": true,
|
| 2062 |
"eos_token": "<|end_of_text|>",
|
| 2063 |
-
"extra_special_tokens": {},
|
| 2064 |
"model_input_names": [
|
| 2065 |
"input_ids",
|
| 2066 |
"attention_mask"
|
|
@@ -2068,5 +2059,5 @@
|
|
| 2068 |
"model_max_length": 15000,
|
| 2069 |
"pad_token": "!",
|
| 2070 |
"padding_side": "right",
|
| 2071 |
-
"tokenizer_class": "
|
| 2072 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"added_tokens_decoder": {
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
"128000": {
|
| 4 |
"content": "<|begin_of_text|>",
|
| 5 |
"lstrip": false,
|
|
|
|
| 2052 |
"bos_token": "<|begin_of_text|>",
|
| 2053 |
"clean_up_tokenization_spaces": true,
|
| 2054 |
"eos_token": "<|end_of_text|>",
|
|
|
|
| 2055 |
"model_input_names": [
|
| 2056 |
"input_ids",
|
| 2057 |
"attention_mask"
|
|
|
|
| 2059 |
"model_max_length": 15000,
|
| 2060 |
"pad_token": "!",
|
| 2061 |
"padding_side": "right",
|
| 2062 |
+
"tokenizer_class": "PreTrainedTokenizerFast"
|
| 2063 |
}
|
trainer_state.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
training_args.bin
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:86bb579b1f9511be627da436115c8d38c5c6117f40dc3da7a05d4b5661cf7a6c
|
| 3 |
+
size 6904
|
zero_to_fp32.py
ADDED
|
@@ -0,0 +1,587 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 5 |
+
|
| 6 |
+
# DeepSpeed Team
|
| 7 |
+
|
| 8 |
+
# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets
|
| 9 |
+
# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in
|
| 10 |
+
# the future. Once extracted, the weights don't require DeepSpeed and can be used in any
|
| 11 |
+
# application.
|
| 12 |
+
#
|
| 13 |
+
# example: python zero_to_fp32.py . pytorch_model.bin
|
| 14 |
+
|
| 15 |
+
import argparse
|
| 16 |
+
import torch
|
| 17 |
+
import glob
|
| 18 |
+
import math
|
| 19 |
+
import os
|
| 20 |
+
import re
|
| 21 |
+
from collections import OrderedDict
|
| 22 |
+
from dataclasses import dataclass
|
| 23 |
+
|
| 24 |
+
# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with
|
| 25 |
+
# DeepSpeed data structures it has to be available in the current python environment.
|
| 26 |
+
from deepspeed.utils import logger
|
| 27 |
+
from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS,
|
| 28 |
+
FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES,
|
| 29 |
+
FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
@dataclass
|
| 33 |
+
class zero_model_state:
|
| 34 |
+
buffers: dict()
|
| 35 |
+
param_shapes: dict()
|
| 36 |
+
shared_params: list
|
| 37 |
+
ds_version: int
|
| 38 |
+
frozen_param_shapes: dict()
|
| 39 |
+
frozen_param_fragments: dict()
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
debug = 0
|
| 43 |
+
|
| 44 |
+
# load to cpu
|
| 45 |
+
device = torch.device('cpu')
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def atoi(text):
|
| 49 |
+
return int(text) if text.isdigit() else text
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def natural_keys(text):
|
| 53 |
+
'''
|
| 54 |
+
alist.sort(key=natural_keys) sorts in human order
|
| 55 |
+
http://nedbatchelder.com/blog/200712/human_sorting.html
|
| 56 |
+
(See Toothy's implementation in the comments)
|
| 57 |
+
'''
|
| 58 |
+
return [atoi(c) for c in re.split(r'(\d+)', text)]
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def get_model_state_file(checkpoint_dir, zero_stage):
|
| 62 |
+
if not os.path.isdir(checkpoint_dir):
|
| 63 |
+
raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist")
|
| 64 |
+
|
| 65 |
+
# there should be only one file
|
| 66 |
+
if zero_stage <= 2:
|
| 67 |
+
file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt")
|
| 68 |
+
elif zero_stage == 3:
|
| 69 |
+
file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt")
|
| 70 |
+
|
| 71 |
+
if not os.path.exists(file):
|
| 72 |
+
raise FileNotFoundError(f"can't find model states file at '{file}'")
|
| 73 |
+
|
| 74 |
+
return file
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def get_checkpoint_files(checkpoint_dir, glob_pattern):
|
| 78 |
+
# XXX: need to test that this simple glob rule works for multi-node setup too
|
| 79 |
+
ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys)
|
| 80 |
+
|
| 81 |
+
if len(ckpt_files) == 0:
|
| 82 |
+
raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'")
|
| 83 |
+
|
| 84 |
+
return ckpt_files
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
def get_optim_files(checkpoint_dir):
|
| 88 |
+
return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt")
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
def get_model_state_files(checkpoint_dir):
|
| 92 |
+
return get_checkpoint_files(checkpoint_dir, "*_model_states.pt")
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def parse_model_states(files):
|
| 96 |
+
zero_model_states = []
|
| 97 |
+
for file in files:
|
| 98 |
+
state_dict = torch.load(file, map_location=device)
|
| 99 |
+
|
| 100 |
+
if BUFFER_NAMES not in state_dict:
|
| 101 |
+
raise ValueError(f"{file} is not a model state checkpoint")
|
| 102 |
+
buffer_names = state_dict[BUFFER_NAMES]
|
| 103 |
+
if debug:
|
| 104 |
+
print("Found buffers:", buffer_names)
|
| 105 |
+
|
| 106 |
+
# recover just the buffers while restoring them to fp32 if they were saved in fp16
|
| 107 |
+
buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names}
|
| 108 |
+
param_shapes = state_dict[PARAM_SHAPES]
|
| 109 |
+
|
| 110 |
+
# collect parameters that are included in param_shapes
|
| 111 |
+
param_names = []
|
| 112 |
+
for s in param_shapes:
|
| 113 |
+
for name in s.keys():
|
| 114 |
+
param_names.append(name)
|
| 115 |
+
|
| 116 |
+
# update with frozen parameters
|
| 117 |
+
frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None)
|
| 118 |
+
if frozen_param_shapes is not None:
|
| 119 |
+
if debug:
|
| 120 |
+
print(f"Found frozen_param_shapes: {frozen_param_shapes}")
|
| 121 |
+
param_names += list(frozen_param_shapes.keys())
|
| 122 |
+
|
| 123 |
+
# handle shared params
|
| 124 |
+
shared_params = [[k, v] for k, v in state_dict["shared_params"].items()]
|
| 125 |
+
|
| 126 |
+
ds_version = state_dict.get(DS_VERSION, None)
|
| 127 |
+
|
| 128 |
+
frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None)
|
| 129 |
+
|
| 130 |
+
z_model_state = zero_model_state(buffers=buffers,
|
| 131 |
+
param_shapes=param_shapes,
|
| 132 |
+
shared_params=shared_params,
|
| 133 |
+
ds_version=ds_version,
|
| 134 |
+
frozen_param_shapes=frozen_param_shapes,
|
| 135 |
+
frozen_param_fragments=frozen_param_fragments)
|
| 136 |
+
zero_model_states.append(z_model_state)
|
| 137 |
+
|
| 138 |
+
return zero_model_states
|
| 139 |
+
|
| 140 |
+
|
| 141 |
+
def parse_optim_states(files, ds_checkpoint_dir):
|
| 142 |
+
|
| 143 |
+
total_files = len(files)
|
| 144 |
+
state_dicts = []
|
| 145 |
+
for f in files:
|
| 146 |
+
state_dict = torch.load(f, map_location=device)
|
| 147 |
+
# immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights
|
| 148 |
+
# and also handle the case where it was already removed by another helper script
|
| 149 |
+
state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None)
|
| 150 |
+
state_dicts.append(state_dict)
|
| 151 |
+
|
| 152 |
+
if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]:
|
| 153 |
+
raise ValueError(f"{files[0]} is not a zero checkpoint")
|
| 154 |
+
zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE]
|
| 155 |
+
world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT]
|
| 156 |
+
|
| 157 |
+
# For ZeRO-2 each param group can have different partition_count as data parallelism for expert
|
| 158 |
+
# parameters can be different from data parallelism for non-expert parameters. So we can just
|
| 159 |
+
# use the max of the partition_count to get the dp world_size.
|
| 160 |
+
|
| 161 |
+
if type(world_size) is list:
|
| 162 |
+
world_size = max(world_size)
|
| 163 |
+
|
| 164 |
+
if world_size != total_files:
|
| 165 |
+
raise ValueError(
|
| 166 |
+
f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. "
|
| 167 |
+
"Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes."
|
| 168 |
+
)
|
| 169 |
+
|
| 170 |
+
# the groups are named differently in each stage
|
| 171 |
+
if zero_stage <= 2:
|
| 172 |
+
fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS
|
| 173 |
+
elif zero_stage == 3:
|
| 174 |
+
fp32_groups_key = FP32_FLAT_GROUPS
|
| 175 |
+
else:
|
| 176 |
+
raise ValueError(f"unknown zero stage {zero_stage}")
|
| 177 |
+
|
| 178 |
+
if zero_stage <= 2:
|
| 179 |
+
fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))]
|
| 180 |
+
elif zero_stage == 3:
|
| 181 |
+
# if there is more than one param group, there will be multiple flattened tensors - one
|
| 182 |
+
# flattened tensor per group - for simplicity merge them into a single tensor
|
| 183 |
+
#
|
| 184 |
+
# XXX: could make the script more memory efficient for when there are multiple groups - it
|
| 185 |
+
# will require matching the sub-lists of param_shapes for each param group flattened tensor
|
| 186 |
+
|
| 187 |
+
fp32_flat_groups = [
|
| 188 |
+
torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts))
|
| 189 |
+
]
|
| 190 |
+
|
| 191 |
+
return zero_stage, world_size, fp32_flat_groups
|
| 192 |
+
|
| 193 |
+
|
| 194 |
+
def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir):
|
| 195 |
+
"""
|
| 196 |
+
Returns fp32 state_dict reconstructed from ds checkpoint
|
| 197 |
+
|
| 198 |
+
Args:
|
| 199 |
+
- ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are)
|
| 200 |
+
|
| 201 |
+
"""
|
| 202 |
+
print(f"Processing zero checkpoint '{ds_checkpoint_dir}'")
|
| 203 |
+
|
| 204 |
+
optim_files = get_optim_files(ds_checkpoint_dir)
|
| 205 |
+
zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir)
|
| 206 |
+
print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}")
|
| 207 |
+
|
| 208 |
+
model_files = get_model_state_files(ds_checkpoint_dir)
|
| 209 |
+
|
| 210 |
+
zero_model_states = parse_model_states(model_files)
|
| 211 |
+
print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}')
|
| 212 |
+
|
| 213 |
+
if zero_stage <= 2:
|
| 214 |
+
return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states)
|
| 215 |
+
elif zero_stage == 3:
|
| 216 |
+
return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states)
|
| 217 |
+
|
| 218 |
+
|
| 219 |
+
def _zero2_merge_frozen_params(state_dict, zero_model_states):
|
| 220 |
+
if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0:
|
| 221 |
+
return
|
| 222 |
+
|
| 223 |
+
frozen_param_shapes = zero_model_states[0].frozen_param_shapes
|
| 224 |
+
frozen_param_fragments = zero_model_states[0].frozen_param_fragments
|
| 225 |
+
|
| 226 |
+
if debug:
|
| 227 |
+
num_elem = sum(s.numel() for s in frozen_param_shapes.values())
|
| 228 |
+
print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}')
|
| 229 |
+
|
| 230 |
+
wanted_params = len(frozen_param_shapes)
|
| 231 |
+
wanted_numel = sum(s.numel() for s in frozen_param_shapes.values())
|
| 232 |
+
avail_numel = sum([p.numel() for p in frozen_param_fragments.values()])
|
| 233 |
+
print(f'Frozen params: Have {avail_numel} numels to process.')
|
| 234 |
+
print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params')
|
| 235 |
+
|
| 236 |
+
total_params = 0
|
| 237 |
+
total_numel = 0
|
| 238 |
+
for name, shape in frozen_param_shapes.items():
|
| 239 |
+
total_params += 1
|
| 240 |
+
unpartitioned_numel = shape.numel()
|
| 241 |
+
total_numel += unpartitioned_numel
|
| 242 |
+
|
| 243 |
+
state_dict[name] = frozen_param_fragments[name]
|
| 244 |
+
|
| 245 |
+
if debug:
|
| 246 |
+
print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ")
|
| 247 |
+
|
| 248 |
+
print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements")
|
| 249 |
+
|
| 250 |
+
|
| 251 |
+
def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states):
|
| 252 |
+
param_shapes = zero_model_states[0].param_shapes
|
| 253 |
+
|
| 254 |
+
# Reconstruction protocol:
|
| 255 |
+
#
|
| 256 |
+
# XXX: document this
|
| 257 |
+
|
| 258 |
+
if debug:
|
| 259 |
+
for i in range(world_size):
|
| 260 |
+
for j in range(len(fp32_flat_groups[0])):
|
| 261 |
+
print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}")
|
| 262 |
+
|
| 263 |
+
# XXX: memory usage doubles here (zero2)
|
| 264 |
+
num_param_groups = len(fp32_flat_groups[0])
|
| 265 |
+
merged_single_partition_of_fp32_groups = []
|
| 266 |
+
for i in range(num_param_groups):
|
| 267 |
+
merged_partitions = [sd[i] for sd in fp32_flat_groups]
|
| 268 |
+
full_single_fp32_vector = torch.cat(merged_partitions, 0)
|
| 269 |
+
merged_single_partition_of_fp32_groups.append(full_single_fp32_vector)
|
| 270 |
+
avail_numel = sum(
|
| 271 |
+
[full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups])
|
| 272 |
+
|
| 273 |
+
if debug:
|
| 274 |
+
wanted_params = sum([len(shapes) for shapes in param_shapes])
|
| 275 |
+
wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes])
|
| 276 |
+
# not asserting if there is a mismatch due to possible padding
|
| 277 |
+
print(f"Have {avail_numel} numels to process.")
|
| 278 |
+
print(f"Need {wanted_numel} numels in {wanted_params} params.")
|
| 279 |
+
|
| 280 |
+
# params
|
| 281 |
+
# XXX: for huge models that can't fit into the host's RAM we will have to recode this to support
|
| 282 |
+
# out-of-core computing solution
|
| 283 |
+
total_numel = 0
|
| 284 |
+
total_params = 0
|
| 285 |
+
for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups):
|
| 286 |
+
offset = 0
|
| 287 |
+
avail_numel = full_single_fp32_vector.numel()
|
| 288 |
+
for name, shape in shapes.items():
|
| 289 |
+
|
| 290 |
+
unpartitioned_numel = shape.numel()
|
| 291 |
+
total_numel += unpartitioned_numel
|
| 292 |
+
total_params += 1
|
| 293 |
+
|
| 294 |
+
if debug:
|
| 295 |
+
print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ")
|
| 296 |
+
state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape)
|
| 297 |
+
offset += unpartitioned_numel
|
| 298 |
+
|
| 299 |
+
# Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and
|
| 300 |
+
# avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex
|
| 301 |
+
# paddings performed in the code it's almost impossible to predict the exact numbers w/o the
|
| 302 |
+
# live optimizer object, so we are checking that the numbers are within the right range
|
| 303 |
+
align_to = 2 * world_size
|
| 304 |
+
|
| 305 |
+
def zero2_align(x):
|
| 306 |
+
return align_to * math.ceil(x / align_to)
|
| 307 |
+
|
| 308 |
+
if debug:
|
| 309 |
+
print(f"original offset={offset}, avail_numel={avail_numel}")
|
| 310 |
+
|
| 311 |
+
offset = zero2_align(offset)
|
| 312 |
+
avail_numel = zero2_align(avail_numel)
|
| 313 |
+
|
| 314 |
+
if debug:
|
| 315 |
+
print(f"aligned offset={offset}, avail_numel={avail_numel}")
|
| 316 |
+
|
| 317 |
+
# Sanity check
|
| 318 |
+
if offset != avail_numel:
|
| 319 |
+
raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong")
|
| 320 |
+
|
| 321 |
+
print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements")
|
| 322 |
+
|
| 323 |
+
|
| 324 |
+
def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states):
|
| 325 |
+
state_dict = OrderedDict()
|
| 326 |
+
|
| 327 |
+
# buffers
|
| 328 |
+
buffers = zero_model_states[0].buffers
|
| 329 |
+
state_dict.update(buffers)
|
| 330 |
+
if debug:
|
| 331 |
+
print(f"added {len(buffers)} buffers")
|
| 332 |
+
|
| 333 |
+
_zero2_merge_frozen_params(state_dict, zero_model_states)
|
| 334 |
+
|
| 335 |
+
_zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states)
|
| 336 |
+
|
| 337 |
+
# recover shared parameters
|
| 338 |
+
for pair in zero_model_states[0].shared_params:
|
| 339 |
+
if pair[1] in state_dict:
|
| 340 |
+
state_dict[pair[0]] = state_dict[pair[1]]
|
| 341 |
+
|
| 342 |
+
return state_dict
|
| 343 |
+
|
| 344 |
+
|
| 345 |
+
def zero3_partitioned_param_info(unpartitioned_numel, world_size):
|
| 346 |
+
remainder = unpartitioned_numel % world_size
|
| 347 |
+
padding_numel = (world_size - remainder) if remainder else 0
|
| 348 |
+
partitioned_numel = math.ceil(unpartitioned_numel / world_size)
|
| 349 |
+
return partitioned_numel, padding_numel
|
| 350 |
+
|
| 351 |
+
|
| 352 |
+
def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states):
|
| 353 |
+
if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0:
|
| 354 |
+
return
|
| 355 |
+
|
| 356 |
+
if debug:
|
| 357 |
+
for i in range(world_size):
|
| 358 |
+
num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values())
|
| 359 |
+
print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}')
|
| 360 |
+
|
| 361 |
+
frozen_param_shapes = zero_model_states[0].frozen_param_shapes
|
| 362 |
+
wanted_params = len(frozen_param_shapes)
|
| 363 |
+
wanted_numel = sum(s.numel() for s in frozen_param_shapes.values())
|
| 364 |
+
avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size
|
| 365 |
+
print(f'Frozen params: Have {avail_numel} numels to process.')
|
| 366 |
+
print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params')
|
| 367 |
+
|
| 368 |
+
total_params = 0
|
| 369 |
+
total_numel = 0
|
| 370 |
+
for name, shape in zero_model_states[0].frozen_param_shapes.items():
|
| 371 |
+
total_params += 1
|
| 372 |
+
unpartitioned_numel = shape.numel()
|
| 373 |
+
total_numel += unpartitioned_numel
|
| 374 |
+
|
| 375 |
+
param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states)
|
| 376 |
+
state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape)
|
| 377 |
+
|
| 378 |
+
partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size)
|
| 379 |
+
|
| 380 |
+
if debug:
|
| 381 |
+
print(
|
| 382 |
+
f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}"
|
| 383 |
+
)
|
| 384 |
+
|
| 385 |
+
print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements")
|
| 386 |
+
|
| 387 |
+
|
| 388 |
+
def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states):
|
| 389 |
+
param_shapes = zero_model_states[0].param_shapes
|
| 390 |
+
avail_numel = fp32_flat_groups[0].numel() * world_size
|
| 391 |
+
# Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each
|
| 392 |
+
# param, re-consolidating each param, while dealing with padding if any
|
| 393 |
+
|
| 394 |
+
# merge list of dicts, preserving order
|
| 395 |
+
param_shapes = {k: v for d in param_shapes for k, v in d.items()}
|
| 396 |
+
|
| 397 |
+
if debug:
|
| 398 |
+
for i in range(world_size):
|
| 399 |
+
print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}")
|
| 400 |
+
|
| 401 |
+
wanted_params = len(param_shapes)
|
| 402 |
+
wanted_numel = sum(shape.numel() for shape in param_shapes.values())
|
| 403 |
+
# not asserting if there is a mismatch due to possible padding
|
| 404 |
+
avail_numel = fp32_flat_groups[0].numel() * world_size
|
| 405 |
+
print(f"Trainable params: Have {avail_numel} numels to process.")
|
| 406 |
+
print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.")
|
| 407 |
+
|
| 408 |
+
# params
|
| 409 |
+
# XXX: for huge models that can't fit into the host's RAM we will have to recode this to support
|
| 410 |
+
# out-of-core computing solution
|
| 411 |
+
offset = 0
|
| 412 |
+
total_numel = 0
|
| 413 |
+
total_params = 0
|
| 414 |
+
for name, shape in param_shapes.items():
|
| 415 |
+
|
| 416 |
+
unpartitioned_numel = shape.numel()
|
| 417 |
+
total_numel += unpartitioned_numel
|
| 418 |
+
total_params += 1
|
| 419 |
+
|
| 420 |
+
partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size)
|
| 421 |
+
|
| 422 |
+
if debug:
|
| 423 |
+
print(
|
| 424 |
+
f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}"
|
| 425 |
+
)
|
| 426 |
+
|
| 427 |
+
# XXX: memory usage doubles here
|
| 428 |
+
state_dict[name] = torch.cat(
|
| 429 |
+
tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)),
|
| 430 |
+
0).narrow(0, 0, unpartitioned_numel).view(shape)
|
| 431 |
+
offset += partitioned_numel
|
| 432 |
+
|
| 433 |
+
offset *= world_size
|
| 434 |
+
|
| 435 |
+
# Sanity check
|
| 436 |
+
if offset != avail_numel:
|
| 437 |
+
raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong")
|
| 438 |
+
|
| 439 |
+
print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements")
|
| 440 |
+
|
| 441 |
+
|
| 442 |
+
def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states):
|
| 443 |
+
state_dict = OrderedDict()
|
| 444 |
+
|
| 445 |
+
# buffers
|
| 446 |
+
buffers = zero_model_states[0].buffers
|
| 447 |
+
state_dict.update(buffers)
|
| 448 |
+
if debug:
|
| 449 |
+
print(f"added {len(buffers)} buffers")
|
| 450 |
+
|
| 451 |
+
_zero3_merge_frozen_params(state_dict, world_size, zero_model_states)
|
| 452 |
+
|
| 453 |
+
_zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states)
|
| 454 |
+
|
| 455 |
+
# recover shared parameters
|
| 456 |
+
for pair in zero_model_states[0].shared_params:
|
| 457 |
+
if pair[1] in state_dict:
|
| 458 |
+
state_dict[pair[0]] = state_dict[pair[1]]
|
| 459 |
+
|
| 460 |
+
return state_dict
|
| 461 |
+
|
| 462 |
+
|
| 463 |
+
def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None):
|
| 464 |
+
"""
|
| 465 |
+
Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with
|
| 466 |
+
``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example
|
| 467 |
+
via a model hub.
|
| 468 |
+
|
| 469 |
+
Args:
|
| 470 |
+
- ``checkpoint_dir``: path to the desired checkpoint folder
|
| 471 |
+
- ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14``
|
| 472 |
+
|
| 473 |
+
Returns:
|
| 474 |
+
- pytorch ``state_dict``
|
| 475 |
+
|
| 476 |
+
Note: this approach may not work if your application doesn't have sufficient free CPU memory and
|
| 477 |
+
you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with
|
| 478 |
+
the checkpoint.
|
| 479 |
+
|
| 480 |
+
A typical usage might be ::
|
| 481 |
+
|
| 482 |
+
from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint
|
| 483 |
+
# do the training and checkpoint saving
|
| 484 |
+
state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu
|
| 485 |
+
model = model.cpu() # move to cpu
|
| 486 |
+
model.load_state_dict(state_dict)
|
| 487 |
+
# submit to model hub or save the model to share with others
|
| 488 |
+
|
| 489 |
+
In this example the ``model`` will no longer be usable in the deepspeed context of the same
|
| 490 |
+
application. i.e. you will need to re-initialize the deepspeed engine, since
|
| 491 |
+
``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it.
|
| 492 |
+
|
| 493 |
+
If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead.
|
| 494 |
+
|
| 495 |
+
"""
|
| 496 |
+
if tag is None:
|
| 497 |
+
latest_path = os.path.join(checkpoint_dir, 'latest')
|
| 498 |
+
if os.path.isfile(latest_path):
|
| 499 |
+
with open(latest_path, 'r') as fd:
|
| 500 |
+
tag = fd.read().strip()
|
| 501 |
+
else:
|
| 502 |
+
raise ValueError(f"Unable to find 'latest' file at {latest_path}")
|
| 503 |
+
|
| 504 |
+
ds_checkpoint_dir = os.path.join(checkpoint_dir, tag)
|
| 505 |
+
|
| 506 |
+
if not os.path.isdir(ds_checkpoint_dir):
|
| 507 |
+
raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist")
|
| 508 |
+
|
| 509 |
+
return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir)
|
| 510 |
+
|
| 511 |
+
|
| 512 |
+
def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None):
|
| 513 |
+
"""
|
| 514 |
+
Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be
|
| 515 |
+
loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed.
|
| 516 |
+
|
| 517 |
+
Args:
|
| 518 |
+
- ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``)
|
| 519 |
+
- ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin)
|
| 520 |
+
- ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14``
|
| 521 |
+
"""
|
| 522 |
+
|
| 523 |
+
state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag)
|
| 524 |
+
print(f"Saving fp32 state dict to {output_file}")
|
| 525 |
+
torch.save(state_dict, output_file)
|
| 526 |
+
|
| 527 |
+
|
| 528 |
+
def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None):
|
| 529 |
+
"""
|
| 530 |
+
1. Put the provided model to cpu
|
| 531 |
+
2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict``
|
| 532 |
+
3. Load it into the provided model
|
| 533 |
+
|
| 534 |
+
Args:
|
| 535 |
+
- ``model``: the model object to update
|
| 536 |
+
- ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``)
|
| 537 |
+
- ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14``
|
| 538 |
+
|
| 539 |
+
Returns:
|
| 540 |
+
- ``model`: modified model
|
| 541 |
+
|
| 542 |
+
Make sure you have plenty of CPU memory available before you call this function. If you don't
|
| 543 |
+
have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it
|
| 544 |
+
conveniently placed for you in the checkpoint folder.
|
| 545 |
+
|
| 546 |
+
A typical usage might be ::
|
| 547 |
+
|
| 548 |
+
from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint
|
| 549 |
+
model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir)
|
| 550 |
+
# submit to model hub or save the model to share with others
|
| 551 |
+
|
| 552 |
+
Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context
|
| 553 |
+
of the same application. i.e. you will need to re-initialize the deepspeed engine, since
|
| 554 |
+
``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it.
|
| 555 |
+
|
| 556 |
+
"""
|
| 557 |
+
logger.info(f"Extracting fp32 weights")
|
| 558 |
+
state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag)
|
| 559 |
+
|
| 560 |
+
logger.info(f"Overwriting model with fp32 weights")
|
| 561 |
+
model = model.cpu()
|
| 562 |
+
model.load_state_dict(state_dict, strict=False)
|
| 563 |
+
|
| 564 |
+
return model
|
| 565 |
+
|
| 566 |
+
|
| 567 |
+
if __name__ == "__main__":
|
| 568 |
+
|
| 569 |
+
parser = argparse.ArgumentParser()
|
| 570 |
+
parser.add_argument("checkpoint_dir",
|
| 571 |
+
type=str,
|
| 572 |
+
help="path to the desired checkpoint folder, e.g., path/checkpoint-12")
|
| 573 |
+
parser.add_argument(
|
| 574 |
+
"output_file",
|
| 575 |
+
type=str,
|
| 576 |
+
help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)")
|
| 577 |
+
parser.add_argument("-t",
|
| 578 |
+
"--tag",
|
| 579 |
+
type=str,
|
| 580 |
+
default=None,
|
| 581 |
+
help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1")
|
| 582 |
+
parser.add_argument("-d", "--debug", action='store_true', help="enable debug")
|
| 583 |
+
args = parser.parse_args()
|
| 584 |
+
|
| 585 |
+
debug = args.debug
|
| 586 |
+
|
| 587 |
+
convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, args.output_file, tag=args.tag)
|