From 64f1834e017f33fc0b56770a9fe60f265ff4b6b1 Mon Sep 17 00:00:00 2001 From: Lewis Tunstall Date: Fri, 10 Nov 2023 10:00:05 +0000 Subject: [PATCH] Add config tests --- tests/fixtures/config_dpo_full.yaml | 37 +++++++++++++++++++++++++ tests/fixtures/config_sft_full.yaml | 41 +++++++++++++++++++++++++++ tests/test_configs.py | 43 +++++++++++++++++++++++++++++ 3 files changed, 121 insertions(+) create mode 100644 tests/fixtures/config_dpo_full.yaml create mode 100644 tests/fixtures/config_sft_full.yaml create mode 100644 tests/test_configs.py diff --git a/tests/fixtures/config_dpo_full.yaml b/tests/fixtures/config_dpo_full.yaml new file mode 100644 index 0000000..5110f59 --- /dev/null +++ b/tests/fixtures/config_dpo_full.yaml @@ -0,0 +1,37 @@ +# Model arguments +model_name_or_path: alignment-handbook/zephyr-7b-sft-full + +# Data training arguments +# For definitions, see: src/h4/training/config.py +dataset_mixer: + HuggingFaceH4/ultrafeedback_binarized: 1.0 +dataset_splits: +- train_prefs +- test_prefs +preprocessing_num_workers: 12 + +# DPOTrainer arguments +bf16: true +beta: 0.1 +do_eval: true +evaluation_strategy: steps +eval_steps: 100 +gradient_accumulation_steps: 1 +gradient_checkpointing: true +hub_model_id: zephyr-7b-dpo-full +learning_rate: 5.0e-7 +log_level: info +logging_steps: 10 +lr_scheduler_type: linear +max_length: 1024 +max_prompt_length: 512 +num_train_epochs: 3 +optim: rmsprop +output_dir: data/zephyr-7b-dpo-full +per_device_train_batch_size: 8 +per_device_eval_batch_size: 4 +push_to_hub: true +save_strategy: "no" +save_total_limit: null +seed: 42 +warmup_ratio: 0.1 \ No newline at end of file diff --git a/tests/fixtures/config_sft_full.yaml b/tests/fixtures/config_sft_full.yaml new file mode 100644 index 0000000..81720e9 --- /dev/null +++ b/tests/fixtures/config_sft_full.yaml @@ -0,0 +1,41 @@ +# Model arguments +model_name_or_path: mistralai/Mistral-7B-v0.1 +model_revision: main +torch_dtype: bfloat16 +use_flash_attention_2: true + +# Data training arguments +dataset_mixer: + HuggingFaceH4/ultrachat_200k: 1.0 +dataset_splits: +- train_sft +- test_sft +preprocessing_num_workers: 12 + +# SFT trainer config +bf16: true +do_eval: true +evaluation_strategy: epoch +gradient_accumulation_steps: 2 +gradient_checkpointing: true +hub_model_id: zephyr-7b-sft-full +hub_strategy: every_save +learning_rate: 2.0e-05 +log_level: info +logging_steps: 5 +logging_strategy: steps +lr_scheduler_type: cosine +max_seq_length: 2048 +max_steps: -1 +num_train_epochs: 1 +output_dir: data/zephyr-7b-sft-full +overwrite_output_dir: true +per_device_eval_batch_size: 16 +per_device_train_batch_size: 32 +push_to_hub: True +remove_unused_columns: true +report_to: +- tensorboard +save_strategy: "no" +save_total_limit: null +seed: 42 \ No newline at end of file diff --git a/tests/test_configs.py b/tests/test_configs.py new file mode 100644 index 0000000..2a4a7a6 --- /dev/null +++ b/tests/test_configs.py @@ -0,0 +1,43 @@ +# Copyright 2023 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os +import unittest + +from alignment import DataArguments, H4ArgumentParser, ModelArguments, SFTConfig + + +class H4ArgumentParserTest(unittest.TestCase): + def setUp(self): + self.parser = H4ArgumentParser((ModelArguments, DataArguments, SFTConfig)) + self.yaml_file_path = "tests/fixtures/config_sft_full.yaml" + + def test_load_yaml(self): + model_args, data_args, training_args = self.parser.parse_yaml_file(os.path.abspath(self.yaml_file_path)) + self.assertEqual(model_args.model_name_or_path, "mistralai/Mistral-7B-v0.1") + + def test_load_yaml_and_args(self): + command_line_args = [ + "--model_name_or_path=test", + "--use_peft=true", + "--lora_r=16", + "--lora_dropout=0.5", + ] + model_args, data_args, training_args = self.parser.parse_yaml_and_args( + os.path.abspath(self.yaml_file_path), command_line_args + ) + self.assertEqual(model_args.model_name_or_path, "test") + self.assertEqual(model_args.use_peft, True) + self.assertEqual(model_args.lora_r, 16) + self.assertEqual(model_args.lora_dropout, 0.5)