已合并
[pytorch][refact]add coverage test cases #3810
HANHU1CHEN创建于 2025年11月28日
[pytorch][refact]add coverage test cases #3810
已合并
从已删除 :master合入到Ascend/MindSpeed-LLMmaster
共 9 个文件变更+304-71
| @@ -632,73 +632,6 @@ def _add_or_replace_eos_token(tokenizer: "PreTrainedTokenizer", eos_token: str) | |||
| 632 | logger.warning("New tokens have been added, make sure `resize_vocab` is True.") | 632 | logger.warning("New tokens have been added, make sure `resize_vocab` is True.") |
| 633 | 633 | ||
| 634 | 634 | ||
| 635 | -def _jinja_escape(content: str) -> str: | ||
| 636 | - return content.replace("\n", r"\n").replace("'", r"\'") | ||
| 637 | - | ||
| 638 | - | ||
| 639 | -def _convert_slots_to_jinja(slots: "SLOTS", tokenizer: "PreTrainedTokenizer", placeholder: str = "content") -> str: | ||
| 640 | - slot_items = [] | ||
| 641 | - for slot in slots: | ||
| 642 | - if isinstance(slot, str): | ||
| 643 | - slot_pieces = slot.split("{{content}}") | ||
| 644 | - if slot_pieces[0]: | ||
| 645 | - slot_items.append("'" + _jinja_escape(slot_pieces[0]) + "'") | ||
| 646 | - if len(slot_pieces) > 1: | ||
| 647 | - slot_items.append(placeholder) | ||
| 648 | - if slot_pieces[1]: | ||
| 649 | - slot_items.append("'" + _jinja_escape(slot_pieces[1]) + "'") | ||
| 650 | - elif isinstance(slot, set): | ||
| 651 | - if "bos_token" in slot: | ||
| 652 | - slot_items.append("'" + tokenizer.bos_token + "'") | ||
| 653 | - elif "eos_token" in slot: # do not use {{ eos_token }} since it may be replaced | ||
| 654 | - slot_items.append("'" + tokenizer.eos_token + "'") | ||
| 655 | - elif isinstance(slot, dict): | ||
| 656 | - raise ValueError("Dict is not supported.") | ||
| 657 | - | ||
| 658 | - return " + ".join(slot_items) | ||
| 659 | - | ||
| 660 | - | ||
| 661 | -def _get_jinja_template(template: "Template", tokenizer: "PreTrainedTokenizer") -> str: | ||
| 662 | - jinja_template = "" | ||
| 663 | - | ||
| 664 | - prefix = _convert_slots_to_jinja(template.format_prefix.apply(), tokenizer) | ||
| 665 | - if prefix: | ||
| 666 | - jinja_template += "{{ " + prefix + " }}" | ||
| 667 | - | ||
| 668 | - if template.default_system: | ||
| 669 | - jinja_template += "{% set system_message = '" + _jinja_escape(template.default_system) + "' %}" | ||
| 670 | - | ||
| 671 | - jinja_template += ( | ||
| 672 | - "{% if messages[0]['role'] == 'system' %}{% set system_message = messages[0]['content'] %}{% endif %}" | ||
| 673 | - ) | ||
| 674 | - | ||
| 675 | - system_message = _convert_slots_to_jinja(template.format_system.apply(), tokenizer, placeholder="system_message") | ||
| 676 | - if isinstance(template, Llama2Template): | ||
| 677 | - pass | ||
| 678 | - elif template.force_system: | ||
| 679 | - jinja_template += "{{ " + system_message + " }}" | ||
| 680 | - else: | ||
| 681 | - jinja_template += "{% if system_message is defined %}{{ " + system_message + " }}{% endif %}" | ||
| 682 | - | ||
| 683 | - jinja_template += "{% for message in messages %}" | ||
| 684 | - jinja_template += "{% set content = message['content'] %}" | ||
| 685 | - if isinstance(template, Llama2Template): | ||
| 686 | - jinja_template += "{% if loop.index0 == 0 and system_message is defined %}" | ||
| 687 | - jinja_template += "{% set content = " + system_message + " + message['content'] %}" | ||
| 688 | - jinja_template += "{% endif %}" | ||
| 689 | - jinja_template += "{% if message['role'] == 'user' %}" | ||
| 690 | - user_message = _convert_slots_to_jinja(template.format_user.apply(), tokenizer) | ||
| 691 | - jinja_template += "{{ " + user_message + " }}" | ||
| 692 | - jinja_template += "{% elif message['role'] == 'assistant' %}" | ||
| 693 | - assistant_message = _convert_slots_to_jinja( | ||
| 694 | - template.format_assistant.apply() + template.format_separator.apply(), tokenizer | ||
| 695 | - ) | ||
| 696 | - jinja_template += "{{ " + assistant_message + " }}" | ||
| 697 | - jinja_template += "{% endif %}" | ||
| 698 | - jinja_template += "{% endfor %}" | ||
| 699 | - return jinja_template | ||
| 700 | - | ||
| 701 | - | ||
| 702 | def register_custom_template(name, json_file_path=TEMPLATES_DIR, enable_thinking=False) -> str: | 635 | def register_custom_template(name, json_file_path=TEMPLATES_DIR, enable_thinking=False) -> str: |
| 703 | if name in templates: | 636 | if name in templates: |
| 704 | return name | 637 | return name |
| @@ -14,7 +14,7 @@ | |||
| 14 | <th>Mem.</th> | 14 | <th>Mem.</th> |
| 15 | </tr> | 15 | </tr> |
| 16 | <tr> | 16 | <tr> |
| 17 | - <td rowspan="15">ST</td> | 17 | + <td rowspan="16">ST</td> |
| 18 | <td rowspan="14">Pretrain</td> | 18 | <td rowspan="14">Pretrain</td> |
| 19 | <td>TP,PP,VPP,distributed_optimizer,o2_gradient,o2_optimizer,重计算,enable_recompute_layers_per_pp_rank,FA_TND,use_fused_rotary_pos_emb</td> | 19 | <td>TP,PP,VPP,distributed_optimizer,o2_gradient,o2_optimizer,重计算,enable_recompute_layers_per_pp_rank,FA_TND,use_fused_rotary_pos_emb</td> |
| 20 | <td><a href="st/shell_scripts/llama2_tp2_pp4_vpp2_ptd.sh">llama2_tp2_pp4_vpp2_ptd.sh</a></td> | 20 | <td><a href="st/shell_scripts/llama2_tp2_pp4_vpp2_ptd.sh">llama2_tp2_pp4_vpp2_ptd.sh</a></td> |
| @@ -121,6 +121,14 @@ | |||
| 121 | <td>Y</td> | 121 | <td>Y</td> |
| 122 | <td>Y</td> | 122 | <td>Y</td> |
| 123 | </tr> | 123 | </tr> |
| 124 | + <tr> | ||
| 125 | + <td rowspan="1">DPO</td> | ||
| 126 | + <td>is_pairwise_dataset, cyclic</td> | ||
| 127 | + <td><a href="st/shell_scripts/dpo_llama2_tp1_pp1_cyclic_pairwise.sh">dpo_llama2_tp1_pp1_cyclic_pairwise.sh</a></td> | ||
| 128 | + <td>Y</td> | ||
| 129 | + <td>Y</td> | ||
| 130 | + <td>Y</td> | ||
| 131 | + </tr> | ||
| 124 | <tr> | 132 | <tr> |
| 125 | <td rowspan="5">UT</td> | 133 | <td rowspan="5">UT</td> |
| 126 | <td>Inference</td> | 134 | <td>Inference</td> |
| @@ -248,7 +256,7 @@ | |||
| 248 | <td></td> | 256 | <td></td> |
| 249 | </tr> | 257 | </tr> |
| 250 | <tr> | 258 | <tr> |
| 251 | - <td rowspan="3"><a href="pipeline/common">ProcessData</td> | 259 | + <td rowspan="5"><a href="pipeline/common">ProcessData</td> |
| 252 | <td>instruction_data_alpaca, | 260 | <td>instruction_data_alpaca, |
| 253 | instruction_data_alpaca_history, | 261 | instruction_data_alpaca_history, |
| 254 | instruction_data_sharegpt, | 262 | instruction_data_sharegpt, |
| @@ -272,6 +280,20 @@ | |||
| 272 | <td></td> | 280 | <td></td> |
| 273 | <td></td> | 281 | <td></td> |
| 274 | </tr> | 282 | </tr> |
| 283 | + <tr> | ||
| 284 | + <td>GPTSentencePieceTokenizer</td> | ||
| 285 | + <td><a href="coverage/process_data/test_process_pretrain_data.py">test_process_pretrain_data.py</a></td> | ||
| 286 | + <td>Y</td> | ||
| 287 | + <td></td> | ||
| 288 | + <td></td> | ||
| 289 | + </tr> | ||
| 290 | + <tr> | ||
| 291 | + <td>reasoning_template</td> | ||
| 292 | + <td><a href="coverage/process_data/test_process_instruction_data.py">test_process_instruction_data.py</a></td> | ||
| 293 | + <td>Y</td> | ||
| 294 | + <td></td> | ||
| 295 | + <td></td> | ||
| 296 | + </tr> | ||
| 275 | <tr> | 297 | <tr> |
| 276 | <td rowspan="2"><a href="pipeline/baichuan2-13B">Baichuan2-13B</a></td> | 298 | <td rowspan="2"><a href="pipeline/baichuan2-13B">Baichuan2-13B</a></td> |
| 277 | <td>data_process</td> | 299 | <td>data_process</td> |
| @@ -106,5 +106,29 @@ | |||
| 106 | "prompt-type" : "llama2" | 106 | "prompt-type" : "llama2" |
| 107 | } | 107 | } |
| 108 | } | 108 | } |
| 109 | + ], | ||
| 110 | + "template_dir": [ | ||
| 111 | + { | ||
| 112 | + "params" : { | ||
| 113 | + "test-out-template": "/data/ci/cache/process_dataset/test_template/", | ||
| 114 | + "base-out-template": "/data/ci/datasets/processed/base_template/" | ||
| 115 | + } | ||
| 116 | + } | ||
| 117 | + ], | ||
| 118 | + "reasoning_template": [ | ||
| 119 | + { | ||
| 120 | + "params": { | ||
| 121 | + "input": "/data/ci/datasets/origin/train-00000-of-00001-a09b74b3ef9c3b56.parquet", | ||
| 122 | + "tokenizer-type": "PretrainedFromHF", | ||
| 123 | + "handler-name": "AlpacaStyleInstructionHandler", | ||
| 124 | + "output-prefix": "/data/ci/cache/process_dataset/test_template/qwen3_reasoning_template", | ||
| 125 | + "tokenizer-name-or-path": "/data/ci/models/qwen3_next/hf/Qwen3-Next-80B-A3B-hf", | ||
| 126 | + "cache-dir": "/data/ci/cache/process_dataset/tmp/", | ||
| 127 | + "workers": 4, | ||
| 128 | + "log-interval": 1000, | ||
| 129 | + "prompt-type": "qwen3", | ||
| 130 | + "enable-thinking": "true" | ||
| 131 | + } | ||
| 132 | + } | ||
| 109 | ] | 133 | ] |
| 110 | } | 134 | } |
| @@ -146,3 +146,28 @@ class TestProcessInstructionDataMultiHandler: | |||
| 146 | test_file = params["output-prefix"] + end_str | 146 | test_file = params["output-prefix"] + end_str |
| 147 | assert compare_file_md5_same(base_file, test_file) | 147 | assert compare_file_md5_same(base_file, test_file) |
| 148 | 148 | ||
| 149 | + | ||
| 150 | +class TestProcessInstructionDataTemplate: | ||
| 151 | + test_config = create_testconfig(Path(__file__).with_suffix(".json")) | ||
| 152 | + | ||
| 153 | + | ||
| 154 | + [(test_config["template_dir"][0], test_config["reasoning_template"][0])]) | ||
| 155 | + def test_reasoning_template(self, build_args, full_params, params): | ||
| 156 | + # create output dir if it doesn't exist | ||
| 157 | + if not os.path.isdir(full_params["test-out-template"]): | ||
| 158 | + os.makedirs(full_params["test-out-template"]) | ||
| 159 | + | ||
| 160 | + # process instruction dataset | ||
| 161 | + print("\n=============== reasoning template test =============") | ||
| 162 | + main() | ||
| 163 | + | ||
| 164 | + # compare file MD5 hashes | ||
| 165 | + prefix_str = params["output-prefix"].split('/')[-1] | ||
| 166 | + mid_strs = ["packed_attention_mask_document", "packed_input_ids_document", "packed_labels_document"] | ||
| 167 | + end_suffixs = [".bin", ".idx"] | ||
| 168 | + for mid_str in mid_strs: | ||
| 169 | + for end_suffix in end_suffixs: | ||
| 170 | + end_str = "_" + mid_str + end_suffix | ||
| 171 | + base_file = full_params["base-out-template"] + prefix_str + end_str | ||
| 172 | + test_file = params["output-prefix"] + end_str | ||
| 173 | + assert compare_file_md5_same(base_file, test_file) | ||
| @@ -6,7 +6,9 @@ | |||
| 6 | "test-out-part": "/data/ci/cache/process_dataset/test_merge_subs/", | 6 | "test-out-part": "/data/ci/cache/process_dataset/test_merge_subs/", |
| 7 | "base-out-part": "/data/ci/datasets/processed/base_merge_subs/", | 7 | "base-out-part": "/data/ci/datasets/processed/base_merge_subs/", |
| 8 | "test-out-merge": "/data/ci/cache/process_dataset/test_merge/", | 8 | "test-out-merge": "/data/ci/cache/process_dataset/test_merge/", |
| 9 | - "base-out-merge": "/data/ci/datasets/processed/base_merge/" | 9 | + "base-out-merge": "/data/ci/datasets/processed/base_merge/", |
| 10 | + "test-out-tokenizer-type": "/data/ci/cache/process_dataset/test_tokenizer_type/", | ||
| 11 | + "base-out-tokenizer-type": "/data/ci/datasets/processed/base_tokenizer_type/" | ||
| 10 | } | 12 | } |
| 11 | } | 13 | } |
| 12 | ], | 14 | ], |
| @@ -42,5 +44,17 @@ | |||
| 42 | "merge-group-keys": "text_document" | 44 | "merge-group-keys": "text_document" |
| 43 | } | 45 | } |
| 44 | } | 46 | } |
| 47 | + ], | ||
| 48 | + "test_pretrain_datasets_GPTSentencePieceTokenizer": [ | ||
| 49 | + { | ||
| 50 | + "params": { | ||
| 51 | + "input": "/data/ci/datasets/origin/train-00000-of-00001-a09b74b3ef9c3b56.parquet", | ||
| 52 | + "tokenizer-type": "GPTSentencePieceTokenizer", | ||
| 53 | + "output-prefix": "/data/ci/cache/process_dataset/test_tokenizer_type/gptsentencepiece", | ||
| 54 | + "tokenizer-model": "/data/ci/models/mamba2/hf/mamba2-2.7b-hf/mt_nlg_plus_multilingual_ja_zh_the_stack_frac_015_256k.model", | ||
| 55 | + "workers": 4, | ||
| 56 | + "log-interval": 1000 | ||
| 57 | + } | ||
| 58 | + } | ||
| 45 | ] | 59 | ] |
| 46 | } | 60 | } |
| @@ -69,4 +69,24 @@ class TestProcessPretrainData: | |||
| 69 | end_str = prefix_str + end_str | 69 | end_str = prefix_str + end_str |
| 70 | base_file = full_params["base-out-merge"] + end_str | 70 | base_file = full_params["base-out-merge"] + end_str |
| 71 | assert compare_file_md5_same(base_file, test_file) | 71 | assert compare_file_md5_same(base_file, test_file) |
| 72 | - | 72 | + |
| 73 | + | ||
| 74 | + | ||
| 75 | + [(test_config["pretrain_dataset"][0], test_config["test_pretrain_datasets_GPTSentencePieceTokenizer"][0])]) | ||
| 76 | + def test_pretrain_datasets_GPTSentencePieceTokenizer(self, build_args, full_params, params): | ||
| 77 | + # create output dir if it doesn't exist | ||
| 78 | + if not os.path.isdir(full_params["test-out-tokenizer-type"]): | ||
| 79 | + os.makedirs(full_params["test-out-tokenizer-type"]) | ||
| 80 | + | ||
| 81 | + # merge pretrain dataset | ||
| 82 | + print("\n=============== merge pretrain datasets =============") | ||
| 83 | + main() | ||
| 84 | + | ||
| 85 | + # compare file MD5 hashes | ||
| 86 | + prefix_str = params["output-prefix"].split('/')[-1] | ||
| 87 | + end_strs = ["_text_document.bin", "_text_document.idx"] | ||
| 88 | + for end_str in end_strs: | ||
| 89 | + test_file = params["output-prefix"] + end_str | ||
| 90 | + end_str = prefix_str + end_str | ||
| 91 | + base_file = full_params["base-out-tokenizer-type"] + end_str | ||
| 92 | + assert compare_file_md5_same(base_file, test_file) | ||
| @@ -117,4 +117,31 @@ | |||
| 117 | <td>/</td> | 117 | <td>/</td> |
| 118 | <td>/</td> | 118 | <td>/</td> |
| 119 | </tr> | 119 | </tr> |
| 120 | + <tr> | ||
| 121 | + <td>dpo_llama2_tp1_pp1_cyclic_pairwise</td> | ||
| 122 | + <td>/data/ci/models/llama2/hf/llama-2-7b-hf</td> | ||
| 123 | + <td>/data/ci/models/llama2/mg/llama2-7b_2l_tp1pp1</td> | ||
| 124 | + <td>/</td> | ||
| 125 | + <td>/data/ci/datasets/processed/orca/orca_rlhf</td> | ||
| 126 | + <td>/</td> | ||
| 127 | + <td>/</td> | ||
| 128 | + </tr> | ||
| 129 | + <tr> | ||
| 130 | + <td>test_pretrain_datasets_GPTSentencePieceTokenizer</td> | ||
| 131 | + <td>/data/ci/models/mamba2/hf/mamba2-2.7b-hf/mt_nlg_plus_multilingual_ja_zh_the_stack_frac_015_256k.model</td> | ||
| 132 | + <td>/</td> | ||
| 133 | + <td>/data/ci/datasets/origin/train-00000-of-00001-a09b74b3ef9c3b56.parquet</td> | ||
| 134 | + <td>/data/ci/cache/process_dataset/test_tokenizer_type/gptsentencepiece</td> | ||
| 135 | + <td>/</td> | ||
| 136 | + <td>/</td> | ||
| 137 | + </tr> | ||
| 138 | + <tr> | ||
| 139 | + <td>test_reasoning_template</td> | ||
| 140 | + <td>/data/ci/models/qwen3_next/hf/Qwen3-Next-80B-A3B-hf</td> | ||
| 141 | + <td>/</td> | ||
| 142 | + <td>/data/ci/datasets/origin/train-00000-of-00001-a09b74b3ef9c3b56.parquet</td> | ||
| 143 | + <td>/data/ci/cache/process_dataset/test_template/qwen3_reasoning_template</td> | ||
| 144 | + <td>/</td> | ||
| 145 | + <td>/data/ci/cache/process_dataset/tmp</td> | ||
| 146 | + </tr> | ||
| 120 | </table> | 147 | </table> |
| @@ -0,0 +1,61 @@ | |||
| 1 | +{ | ||
| 2 | + "lm loss": [ | ||
| 3 | + 0.6931473, | ||
| 4 | + 0.6931473, | ||
| 5 | + 0.6217661, | ||
| 6 | + 0.583343, | ||
| 7 | + 0.5918418, | ||
| 8 | + 0.4621532, | ||
| 9 | + 0.455303, | ||
| 10 | + 0.5041397, | ||
| 11 | + 0.3615817, | ||
| 12 | + 0.396199, | ||
| 13 | + 0.2921073, | ||
| 14 | + 0.8679615, | ||
| 15 | + 0.6653341, | ||
| 16 | + 0.368182, | ||
| 17 | + 0.7402841 | ||
| 18 | + ], | ||
| 19 | + "grad norm": [ | ||
| 20 | + 639.311, | ||
| 21 | + 603.288, | ||
| 22 | + 685.213, | ||
| 23 | + 294.426, | ||
| 24 | + 222.96, | ||
| 25 | + 197.315, | ||
| 26 | + 91.061, | ||
| 27 | + 103.371, | ||
| 28 | + 95.088, | ||
| 29 | + 73.028, | ||
| 30 | + 108.448, | ||
| 31 | + 337.994, | ||
| 32 | + 233.234, | ||
| 33 | + 57.032, | ||
| 34 | + 286.496 | ||
| 35 | + ], | ||
| 36 | + "time info": [ | ||
| 37 | + 2118.4, | ||
| 38 | + 828.3, | ||
| 39 | + 828.2, | ||
| 40 | + 840.8, | ||
| 41 | + 724.1, | ||
| 42 | + 755.2, | ||
| 43 | + 745.2, | ||
| 44 | + 738.3, | ||
| 45 | + 675.3, | ||
| 46 | + 826.1, | ||
| 47 | + 694.6, | ||
| 48 | + 752.1, | ||
| 49 | + 686.6, | ||
| 50 | + 715.4, | ||
| 51 | + 716.0 | ||
| 52 | + ], | ||
| 53 | + "throughput": [], | ||
| 54 | + "memo info": [ | ||
| 55 | + { | ||
| 56 | + "rank": 0, | ||
| 57 | + "allocated memory": 12723.88330078125, | ||
| 58 | + "max allocated memory": 12739.8984375 | ||
| 59 | + } | ||
| 60 | + ] | ||
| 61 | +} | ||
| @@ -0,0 +1,107 @@ | |||
| 1 | +#!/bin/bash | ||
| 2 | + | ||
| 3 | +export CUDA_DEVICE_MAX_CONNECTIONS=1 | ||
| 4 | +export PYTORCH_NPU_ALLOC_CONF=expandable_segments:True | ||
| 5 | +export HCCL_CONNECT_TIMEOUT=1200 | ||
| 6 | +export HCCL_EXEC_TIMEOUT=1200 | ||
| 7 | + | ||
| 8 | +NPUS_PER_NODE=1 | ||
| 9 | +MASTER_ADDR=localhost | ||
| 10 | +MASTER_PORT=6014 | ||
| 11 | +NNODES=1 | ||
| 12 | +NODE_RANK=0 | ||
| 13 | +WORLD_SIZE=$(($NPUS_PER_NODE*$NNODES)) | ||
| 14 | + | ||
| 15 | +basepath=$(cd `dirname $0`; cd ../../../; pwd) | ||
| 16 | + | ||
| 17 | +CKPT_LOAD_DIR="/data/ci/models/llama2/mg/llama2-7b_2l_tp1pp1/" | ||
| 18 | +DATA_PATH="/data/ci/datasets/processed/orca/orca_rlhf" | ||
| 19 | +TOKENIZER_MODEL="/data/ci/models/llama2/hf/llama-2-7b-hf" | ||
| 20 | + | ||
| 21 | +TP=1 | ||
| 22 | +PP=1 | ||
| 23 | + | ||
| 24 | +DISTRIBUTED_ARGS=( | ||
| 25 | + --nproc_per_node $NPUS_PER_NODE | ||
| 26 | + --nnodes $NNODES | ||
| 27 | + --node_rank $NODE_RANK | ||
| 28 | + --master_addr $MASTER_ADDR | ||
| 29 | + --master_port $MASTER_PORT | ||
| 30 | +) | ||
| 31 | + | ||
| 32 | +GPT_ARGS=( | ||
| 33 | + --no-pad-to-seq-lengths | ||
| 34 | + --use-mcore-models | ||
| 35 | + --tensor-model-parallel-size ${TP} | ||
| 36 | + --pipeline-model-parallel-size ${PP} | ||
| 37 | + --sequence-parallel | ||
| 38 | + --num-layers 2 | ||
| 39 | + --hidden-size 4096 | ||
| 40 | + --ffn-hidden-size 11008 | ||
| 41 | + --num-attention-heads 32 | ||
| 42 | + --tokenizer-type PretrainedFromHF | ||
| 43 | + --tokenizer-name-or-path ${TOKENIZER_MODEL} | ||
| 44 | + --seq-length 4096 | ||
| 45 | + --max-position-embeddings 4096 | ||
| 46 | + --micro-batch-size 1 | ||
| 47 | + --global-batch-size 16 | ||
| 48 | + --make-vocab-size-divisible-by 1 | ||
| 49 | + --lr 1.25e-6 | ||
| 50 | + --train-iters 15 | ||
| 51 | + --lr-decay-style cosine | ||
| 52 | + --untie-embeddings-and-output-weights | ||
| 53 | + --disable-bias-linear | ||
| 54 | + --attention-dropout 0.0 | ||
| 55 | + --init-method-std 0.01 | ||
| 56 | + --hidden-dropout 0.0 | ||
| 57 | + --position-embedding-type rope | ||
| 58 | + --normalization RMSNorm | ||
| 59 | + --use-fused-rmsnorm | ||
| 60 | + --swiglu | ||
| 61 | + --use-flash-attn | ||
| 62 | + --no-masked-softmax-fusion | ||
| 63 | + --attention-softmax-in-fp32 | ||
| 64 | + --min-lr 1.25e-7 | ||
| 65 | + --weight-decay 1e-1 | ||
| 66 | + --lr-warmup-fraction 0.01 | ||
| 67 | + --clip-grad 1.0 | ||
| 68 | + --adam-beta1 0.9 | ||
| 69 | + --initial-loss-scale 65536 | ||
| 70 | + --adam-beta2 0.95 | ||
| 71 | + --no-gradient-accumulation-fusion | ||
| 72 | + --no-load-optim | ||
| 73 | + --no-load-rng | ||
| 74 | + --use-distributed-optimizer | ||
| 75 | + --use-fused-swiglu | ||
| 76 | + --use-fused-rotary-pos-emb | ||
| 77 | + --overlap-grad-reduce | ||
| 78 | + --bf16 | ||
| 79 | +) | ||
| 80 | + | ||
| 81 | +DATA_ARGS=( | ||
| 82 | + --data-path $DATA_PATH | ||
| 83 | + --split 100,0,0 | ||
| 84 | + --log-throughput | ||
| 85 | + --dataloader-type cyclic | ||
| 86 | +) | ||
| 87 | + | ||
| 88 | +RL_ARGS=( | ||
| 89 | + --stage dpo | ||
| 90 | + --dpo-loss-type sigmoid | ||
| 91 | + --is-pairwise-dataset | ||
| 92 | +) | ||
| 93 | + | ||
| 94 | +OUTPUT_ARGS=( | ||
| 95 | + --log-interval 1 | ||
| 96 | + --eval-interval 1000 | ||
| 97 | + --eval-iters 0 | ||
| 98 | +) | ||
| 99 | + | ||
| 100 | +torchrun ${DISTRIBUTED_ARGS[@]} $basepath/posttrain_gpt.py \ | ||
| 101 | + ${GPT_ARGS[@]} \ | ||
| 102 | + ${DATA_ARGS[@]} \ | ||
| 103 | + ${RL_ARGS[@]} \ | ||
| 104 | + ${OUTPUT_ARGS[@]} \ | ||
| 105 | + --load ${CKPT_LOAD_DIR} \ | ||
| 106 | + --finetune \ | ||
| 107 | + --distributed-backend nccl | ||