From 9e25ac1119c668656c89f4f44ae1c20117dda4c4 Mon Sep 17 00:00:00 2001 From: 51616 Date: Fri, 3 Jan 2025 09:31:24 +0000 Subject: [PATCH] trainable on ctx_numbers_128 w/ online ctx_features compute --- configs/context_numbers_10.yaml | 64 ++++ configs/context_numbers_128.yaml | 308 ++++++++++++++++++ configs/context_numbers_debug.yaml | 44 +++ configs/context_numbers_easy.yaml | 53 +++ configs/default.yaml | 12 +- .../context_number_xxlarge/generate_data.py | 79 ----- .../context_numbers_large/generate_data.py | 79 ----- .../context_numbers_medium/generate_data.py | 79 ----- .../context_numbers_small/generate_data.py | 61 ---- .../generate_data.py | 34 +- hyperlora/configs.py | 17 +- hyperlora/data_utils.py | 141 ++++---- hyperlora/intx_sft.py | 221 ++++++++----- hyperlora/modeling_utils.py | 130 +++++--- hyperlora/training_utils.py | 12 +- hyperlora/utils.py | 10 - 16 files changed, 829 insertions(+), 515 deletions(-) create mode 100644 configs/context_numbers_10.yaml create mode 100644 configs/context_numbers_128.yaml create mode 100644 configs/context_numbers_debug.yaml create mode 100644 configs/context_numbers_easy.yaml delete mode 100644 data/raw_datasets/context_number_xxlarge/generate_data.py delete mode 100644 data/raw_datasets/context_numbers_large/generate_data.py delete mode 100644 data/raw_datasets/context_numbers_medium/generate_data.py delete mode 100644 data/raw_datasets/context_numbers_small/generate_data.py rename data/raw_datasets/{context_number_xlarge => }/generate_data.py (66%) diff --git a/configs/context_numbers_10.yaml b/configs/context_numbers_10.yaml new file mode 100644 index 0000000..ba26c4c --- /dev/null +++ b/configs/context_numbers_10.yaml @@ -0,0 +1,64 @@ +output_dir: "" # just a placeholder +bf16: true +model_name_or_path: meta-llama/Llama-3.2-1B-Instruct +label_names: ["labels"] +eval_on_start: True +eval_strategy: "steps" +eval_steps: 500 +save_strategy: "no" +# save_steps: 500 +logging_strategy: "steps" +logging_steps: 100 +use_liger_kernel: true +remove_unused_columns: false + +# needed to avoid OOM by compute the metrics batch by batch +# w/o this the trainer stores logits of all sample in memory... +batch_eval_metrics: true + +per_device_train_batch_size: 128 +per_device_eval_batch_size: 128 +# optim: schedule_free_adamw +learning_rate: 0.00001 +# lr_scheduler_type: "constant_with_warmup" +neftune_noise_alpha: 1 +weight_decay: 0.01 +warmup_ratio: 0.1 + +# LoRA +lora_r: 16 +lora_dropout: 0.05 +target_modules: + - down_proj + - up_proj + - gate_proj + +# data +train_ds_names: +- data/raw_datasets/context_numbers_2 +- data/raw_datasets/context_numbers_3 +- data/raw_datasets/context_numbers_4 +- data/raw_datasets/context_numbers_5 +- data/raw_datasets/context_numbers_6 +- data/raw_datasets/context_numbers_7 +- data/raw_datasets/context_numbers_8 +- data/raw_datasets/context_numbers_9 +- data/raw_datasets/context_numbers_10 + +val_ds_names: +- data/raw_datasets/context_numbers_2 +- data/raw_datasets/context_numbers_3 +- data/raw_datasets/context_numbers_4 +- data/raw_datasets/context_numbers_5 +- data/raw_datasets/context_numbers_6 +- data/raw_datasets/context_numbers_7 +- data/raw_datasets/context_numbers_8 +- data/raw_datasets/context_numbers_9 +- data/raw_datasets/context_numbers_10 + +test_ds_names: +- data/raw_datasets/context_numbers_11 +- data/raw_datasets/context_numbers_12 +- data/raw_datasets/context_numbers_13 +- data/raw_datasets/context_numbers_14 +- data/raw_datasets/context_numbers_15 diff --git a/configs/context_numbers_128.yaml b/configs/context_numbers_128.yaml new file mode 100644 index 0000000..6faf010 --- /dev/null +++ b/configs/context_numbers_128.yaml @@ -0,0 +1,308 @@ +output_dir: "" # just a placeholder +bf16: true +model_name_or_path: meta-llama/Llama-3.2-1B-Instruct +label_names: ["labels"] +eval_on_start: True +eval_strategy: "steps" +eval_steps: 500 +save_strategy: "no" +# save_steps: 500 +logging_strategy: "steps" +logging_steps: 100 +use_liger_kernel: true +remove_unused_columns: false + +# needed to avoid OOM by compute the metrics batch by batch +# w/o this the trainer stores logits of all sample in memory... +batch_eval_metrics: true + +per_device_train_batch_size: 128 +per_device_eval_batch_size: 128 +# optim: schedule_free_adamw +learning_rate: 0.00001 +# lr_scheduler_type: "constant_with_warmup" +neftune_noise_alpha: 1 +weight_decay: 0.01 +warmup_ratio: 0.1 + +# LoRA +lora_r: 16 +lora_dropout: 0.05 +target_modules: + - down_proj + - up_proj + - gate_proj + +# data +train_ds_names: +- data/raw_datasets/context_numbers_2 +- data/raw_datasets/context_numbers_3 +- data/raw_datasets/context_numbers_4 +- data/raw_datasets/context_numbers_5 +- data/raw_datasets/context_numbers_6 +- data/raw_datasets/context_numbers_7 +- data/raw_datasets/context_numbers_8 +- data/raw_datasets/context_numbers_9 +- data/raw_datasets/context_numbers_10 +- data/raw_datasets/context_numbers_11 +- data/raw_datasets/context_numbers_12 +- data/raw_datasets/context_numbers_13 +- data/raw_datasets/context_numbers_14 +- data/raw_datasets/context_numbers_15 +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_17 +- data/raw_datasets/context_numbers_18 +- data/raw_datasets/context_numbers_19 +- data/raw_datasets/context_numbers_20 +- data/raw_datasets/context_numbers_21 +- data/raw_datasets/context_numbers_22 +- data/raw_datasets/context_numbers_23 +- data/raw_datasets/context_numbers_24 +- data/raw_datasets/context_numbers_25 +- data/raw_datasets/context_numbers_26 +- data/raw_datasets/context_numbers_27 +- data/raw_datasets/context_numbers_28 +- data/raw_datasets/context_numbers_29 +- data/raw_datasets/context_numbers_30 +- data/raw_datasets/context_numbers_31 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_33 +- data/raw_datasets/context_numbers_34 +- data/raw_datasets/context_numbers_35 +- data/raw_datasets/context_numbers_36 +- data/raw_datasets/context_numbers_37 +- data/raw_datasets/context_numbers_38 +- data/raw_datasets/context_numbers_39 +- data/raw_datasets/context_numbers_40 +- data/raw_datasets/context_numbers_41 +- data/raw_datasets/context_numbers_42 +- data/raw_datasets/context_numbers_43 +- data/raw_datasets/context_numbers_44 +- data/raw_datasets/context_numbers_45 +- data/raw_datasets/context_numbers_46 +- data/raw_datasets/context_numbers_47 +- data/raw_datasets/context_numbers_48 +- data/raw_datasets/context_numbers_49 +- data/raw_datasets/context_numbers_50 +- data/raw_datasets/context_numbers_51 +- data/raw_datasets/context_numbers_52 +- data/raw_datasets/context_numbers_53 +- data/raw_datasets/context_numbers_54 +- data/raw_datasets/context_numbers_55 +- data/raw_datasets/context_numbers_56 +- data/raw_datasets/context_numbers_57 +- data/raw_datasets/context_numbers_58 +- data/raw_datasets/context_numbers_59 +- data/raw_datasets/context_numbers_60 +- data/raw_datasets/context_numbers_61 +- data/raw_datasets/context_numbers_62 +- data/raw_datasets/context_numbers_63 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_65 +- data/raw_datasets/context_numbers_66 +- data/raw_datasets/context_numbers_67 +- data/raw_datasets/context_numbers_68 +- data/raw_datasets/context_numbers_69 +- data/raw_datasets/context_numbers_70 +- data/raw_datasets/context_numbers_71 +- data/raw_datasets/context_numbers_72 +- data/raw_datasets/context_numbers_73 +- data/raw_datasets/context_numbers_74 +- data/raw_datasets/context_numbers_75 +- data/raw_datasets/context_numbers_76 +- data/raw_datasets/context_numbers_77 +- data/raw_datasets/context_numbers_78 +- data/raw_datasets/context_numbers_79 +- data/raw_datasets/context_numbers_80 +- data/raw_datasets/context_numbers_81 +- data/raw_datasets/context_numbers_82 +- data/raw_datasets/context_numbers_83 +- data/raw_datasets/context_numbers_84 +- data/raw_datasets/context_numbers_85 +- data/raw_datasets/context_numbers_86 +- data/raw_datasets/context_numbers_87 +- data/raw_datasets/context_numbers_88 +- data/raw_datasets/context_numbers_89 +- data/raw_datasets/context_numbers_90 +- data/raw_datasets/context_numbers_91 +- data/raw_datasets/context_numbers_92 +- data/raw_datasets/context_numbers_93 +- data/raw_datasets/context_numbers_94 +- data/raw_datasets/context_numbers_95 +- data/raw_datasets/context_numbers_96 +- data/raw_datasets/context_numbers_97 +- data/raw_datasets/context_numbers_98 +- data/raw_datasets/context_numbers_99 +- data/raw_datasets/context_numbers_100 +- data/raw_datasets/context_numbers_101 +- data/raw_datasets/context_numbers_102 +- data/raw_datasets/context_numbers_103 +- data/raw_datasets/context_numbers_104 +- data/raw_datasets/context_numbers_105 +- data/raw_datasets/context_numbers_106 +- data/raw_datasets/context_numbers_107 +- data/raw_datasets/context_numbers_108 +- data/raw_datasets/context_numbers_109 +- data/raw_datasets/context_numbers_110 +- data/raw_datasets/context_numbers_111 +- data/raw_datasets/context_numbers_112 +- data/raw_datasets/context_numbers_113 +- data/raw_datasets/context_numbers_114 +- data/raw_datasets/context_numbers_115 +- data/raw_datasets/context_numbers_116 +- data/raw_datasets/context_numbers_117 +- data/raw_datasets/context_numbers_118 +- data/raw_datasets/context_numbers_119 +- data/raw_datasets/context_numbers_120 +- data/raw_datasets/context_numbers_121 +- data/raw_datasets/context_numbers_122 +- data/raw_datasets/context_numbers_123 +- data/raw_datasets/context_numbers_124 +- data/raw_datasets/context_numbers_125 +- data/raw_datasets/context_numbers_126 +- data/raw_datasets/context_numbers_127 +- data/raw_datasets/context_numbers_128 + + +val_ds_names: +- data/raw_datasets/context_numbers_2 +- data/raw_datasets/context_numbers_3 +- data/raw_datasets/context_numbers_4 +- data/raw_datasets/context_numbers_5 +- data/raw_datasets/context_numbers_6 +- data/raw_datasets/context_numbers_7 +- data/raw_datasets/context_numbers_8 +- data/raw_datasets/context_numbers_9 +- data/raw_datasets/context_numbers_10 +- data/raw_datasets/context_numbers_11 +- data/raw_datasets/context_numbers_12 +- data/raw_datasets/context_numbers_13 +- data/raw_datasets/context_numbers_14 +- data/raw_datasets/context_numbers_15 +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_17 +- data/raw_datasets/context_numbers_18 +- data/raw_datasets/context_numbers_19 +- data/raw_datasets/context_numbers_20 +- data/raw_datasets/context_numbers_21 +- data/raw_datasets/context_numbers_22 +- data/raw_datasets/context_numbers_23 +- data/raw_datasets/context_numbers_24 +- data/raw_datasets/context_numbers_25 +- data/raw_datasets/context_numbers_26 +- data/raw_datasets/context_numbers_27 +- data/raw_datasets/context_numbers_28 +- data/raw_datasets/context_numbers_29 +- data/raw_datasets/context_numbers_30 +- data/raw_datasets/context_numbers_31 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_33 +- data/raw_datasets/context_numbers_34 +- data/raw_datasets/context_numbers_35 +- data/raw_datasets/context_numbers_36 +- data/raw_datasets/context_numbers_37 +- data/raw_datasets/context_numbers_38 +- data/raw_datasets/context_numbers_39 +- data/raw_datasets/context_numbers_40 +- data/raw_datasets/context_numbers_41 +- data/raw_datasets/context_numbers_42 +- data/raw_datasets/context_numbers_43 +- data/raw_datasets/context_numbers_44 +- data/raw_datasets/context_numbers_45 +- data/raw_datasets/context_numbers_46 +- data/raw_datasets/context_numbers_47 +- data/raw_datasets/context_numbers_48 +- data/raw_datasets/context_numbers_49 +- data/raw_datasets/context_numbers_50 +- data/raw_datasets/context_numbers_51 +- data/raw_datasets/context_numbers_52 +- data/raw_datasets/context_numbers_53 +- data/raw_datasets/context_numbers_54 +- data/raw_datasets/context_numbers_55 +- data/raw_datasets/context_numbers_56 +- data/raw_datasets/context_numbers_57 +- data/raw_datasets/context_numbers_58 +- data/raw_datasets/context_numbers_59 +- data/raw_datasets/context_numbers_60 +- data/raw_datasets/context_numbers_61 +- data/raw_datasets/context_numbers_62 +- data/raw_datasets/context_numbers_63 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_65 +- data/raw_datasets/context_numbers_66 +- data/raw_datasets/context_numbers_67 +- data/raw_datasets/context_numbers_68 +- data/raw_datasets/context_numbers_69 +- data/raw_datasets/context_numbers_70 +- data/raw_datasets/context_numbers_71 +- data/raw_datasets/context_numbers_72 +- data/raw_datasets/context_numbers_73 +- data/raw_datasets/context_numbers_74 +- data/raw_datasets/context_numbers_75 +- data/raw_datasets/context_numbers_76 +- data/raw_datasets/context_numbers_77 +- data/raw_datasets/context_numbers_78 +- data/raw_datasets/context_numbers_79 +- data/raw_datasets/context_numbers_80 +- data/raw_datasets/context_numbers_81 +- data/raw_datasets/context_numbers_82 +- data/raw_datasets/context_numbers_83 +- data/raw_datasets/context_numbers_84 +- data/raw_datasets/context_numbers_85 +- data/raw_datasets/context_numbers_86 +- data/raw_datasets/context_numbers_87 +- data/raw_datasets/context_numbers_88 +- data/raw_datasets/context_numbers_89 +- data/raw_datasets/context_numbers_90 +- data/raw_datasets/context_numbers_91 +- data/raw_datasets/context_numbers_92 +- data/raw_datasets/context_numbers_93 +- data/raw_datasets/context_numbers_94 +- data/raw_datasets/context_numbers_95 +- data/raw_datasets/context_numbers_96 +- data/raw_datasets/context_numbers_97 +- data/raw_datasets/context_numbers_98 +- data/raw_datasets/context_numbers_99 +- data/raw_datasets/context_numbers_100 +- data/raw_datasets/context_numbers_101 +- data/raw_datasets/context_numbers_102 +- data/raw_datasets/context_numbers_103 +- data/raw_datasets/context_numbers_104 +- data/raw_datasets/context_numbers_105 +- data/raw_datasets/context_numbers_106 +- data/raw_datasets/context_numbers_107 +- data/raw_datasets/context_numbers_108 +- data/raw_datasets/context_numbers_109 +- data/raw_datasets/context_numbers_110 +- data/raw_datasets/context_numbers_111 +- data/raw_datasets/context_numbers_112 +- data/raw_datasets/context_numbers_113 +- data/raw_datasets/context_numbers_114 +- data/raw_datasets/context_numbers_115 +- data/raw_datasets/context_numbers_116 +- data/raw_datasets/context_numbers_117 +- data/raw_datasets/context_numbers_118 +- data/raw_datasets/context_numbers_119 +- data/raw_datasets/context_numbers_120 +- data/raw_datasets/context_numbers_121 +- data/raw_datasets/context_numbers_122 +- data/raw_datasets/context_numbers_123 +- data/raw_datasets/context_numbers_124 +- data/raw_datasets/context_numbers_125 +- data/raw_datasets/context_numbers_126 +- data/raw_datasets/context_numbers_127 +- data/raw_datasets/context_numbers_128 + +test_ds_names: +- data/raw_datasets/context_numbers_144 +- data/raw_datasets/context_numbers_160 +- data/raw_datasets/context_numbers_176 +- data/raw_datasets/context_numbers_192 +- data/raw_datasets/context_numbers_208 +- data/raw_datasets/context_numbers_224 +- data/raw_datasets/context_numbers_240 +- data/raw_datasets/context_numbers_256 +- data/raw_datasets/context_numbers_512 +- data/raw_datasets/context_numbers_1024 +- data/raw_datasets/context_numbers_2048 + diff --git a/configs/context_numbers_debug.yaml b/configs/context_numbers_debug.yaml new file mode 100644 index 0000000..70f315a --- /dev/null +++ b/configs/context_numbers_debug.yaml @@ -0,0 +1,44 @@ +output_dir: "" # just a placeholder +bf16: true +model_name_or_path: meta-llama/Llama-3.2-1B-Instruct +label_names: ["labels"] +eval_on_start: True +eval_strategy: "steps" +eval_steps: 500 +save_strategy: "no" +# save_steps: 500 +logging_strategy: "steps" +logging_steps: 100 +use_liger_kernel: true +remove_unused_columns: false + +# needed to avoid OOM by compute the metrics batch by batch +# w/o this the trainer stores logits of all sample in memory... +batch_eval_metrics: true + +per_device_train_batch_size: 128 +per_device_eval_batch_size: 128 +# optim: schedule_free_adamw +learning_rate: 0.00001 +# lr_scheduler_type: "constant_with_warmup" +neftune_noise_alpha: 1 +weight_decay: 0.01 +warmup_ratio: 0.1 + +# LoRA +lora_r: 16 +lora_dropout: 0.05 +target_modules: + - down_proj + - up_proj + - gate_proj + +# data +train_ds_names: +- data/raw_datasets/context_numbers_2 + +val_ds_names: +- data/raw_datasets/context_numbers_2 + +test_ds_names: +- data/raw_datasets/context_numbers_2 diff --git a/configs/context_numbers_easy.yaml b/configs/context_numbers_easy.yaml new file mode 100644 index 0000000..6d8c60c --- /dev/null +++ b/configs/context_numbers_easy.yaml @@ -0,0 +1,53 @@ +output_dir: "" # just a placeholder +bf16: true +model_name_or_path: meta-llama/Llama-3.2-1B-Instruct +label_names: ["labels"] +eval_on_start: True +eval_strategy: "steps" +eval_steps: 500 +save_strategy: "no" +# save_steps: 500 +logging_strategy: "steps" +logging_steps: 100 +use_liger_kernel: true +remove_unused_columns: false + +# needed to avoid OOM by compute the metrics batch by batch +# w/o this the trainer stores logits of all sample in memory... +batch_eval_metrics: true + +per_device_train_batch_size: 128 +per_device_eval_batch_size: 128 +# optim: schedule_free_adamw +learning_rate: 0.00001 +# lr_scheduler_type: "constant_with_warmup" +neftune_noise_alpha: 1 +weight_decay: 0.01 +warmup_ratio: 0.1 + +# LoRA +lora_r: 16 +lora_dropout: 0.05 +target_modules: + - down_proj + - up_proj + - gate_proj + +# data +train_ds_names: +- data/raw_datasets/context_numbers_2 +- data/raw_datasets/context_numbers_3 +- data/raw_datasets/context_numbers_4 +- data/raw_datasets/context_numbers_5 + +val_ds_names: +- data/raw_datasets/context_numbers_2 +- data/raw_datasets/context_numbers_3 +- data/raw_datasets/context_numbers_4 +- data/raw_datasets/context_numbers_5 + +test_ds_names: +- data/raw_datasets/context_numbers_2 +- data/raw_datasets/context_numbers_3 +- data/raw_datasets/context_numbers_4 +- data/raw_datasets/context_numbers_5 diff --git a/configs/default.yaml b/configs/default.yaml index e7dc8d6..15f3035 100644 --- a/configs/default.yaml +++ b/configs/default.yaml @@ -18,19 +18,17 @@ batch_eval_metrics: true per_device_train_batch_size: 128 per_device_eval_batch_size: 128 -optim: schedule_free_adamw -learning_rate: 0.0001 -lr_scheduler_type: "constant_with_warmup" -neftune_noise_alpha: 0 +# optim: schedule_free_adamw +learning_rate: 0.00001 +# lr_scheduler_type: "constant_with_warmup" +neftune_noise_alpha: 1 weight_decay: 0.01 warmup_ratio: 0.1 - - # LoRA lora_r: 8 lora_dropout: 0.05 target_modules: - down_proj - up_proj - - gate_proj \ No newline at end of file + - gate_proj diff --git a/data/raw_datasets/context_number_xxlarge/generate_data.py b/data/raw_datasets/context_number_xxlarge/generate_data.py deleted file mode 100644 index 9b66444..0000000 --- a/data/raw_datasets/context_number_xxlarge/generate_data.py +++ /dev/null @@ -1,79 +0,0 @@ -import itertools -import json -import os -import random -from typing import Dict, List - - -def save_jsonl(data: list[dict], filepath: str) -> None: - """Save data to a JSONL file.""" - parent_dir = os.path.dirname(filepath) - if parent_dir: # Only create directories if there's a parent path - os.makedirs(parent_dir, exist_ok=True) - with open(filepath, "w") as f: - for entry in data: - json.dump(entry, f) - f.write("\n") - - -def get_random_combinations(numbers: list[int], n: int, k: int) -> list[list[int]]: - """Get n random combinations of k numbers from the list.""" - # using itertools.combinations hangs with large numbers - # so we explicitly generate the n indices - indices = [random.choices(range(len(numbers)), k=n) for _ in range(k)] - return zip(*[[numbers[i] for i in ind] for ind in indices]) - - -def generate_number_dataset(max_num: int = 1000, k: int = 3): - """ - Generate a dataset of numbers with corresponding query and answer, - split into train/val/test sets. - - Args: - max_num: Maximum number in the range (exclusive) - k: Number of elements in each combination - """ - # Generate list of all numbers and shuffle them - numbers = list(range(max_num)) - random.shuffle(numbers) - - # Create dataset entries - dataset = [] - query = f"What's your top-{k} favourite numbers in [0-999]? Answer with only the numbers separated by commas." - - # Generate all unique combinations of k numbers - combinations = get_random_combinations(numbers, 120_000, k) - # random.shuffle(combinations) - - for combination in combinations: - entry = { - "context": ", ".join(map(str, combination)), - "prompt": query, - "response": ", ".join(map(str, combination)), - } - dataset.append(entry) - - # Calculate split sizes - total_size = len(dataset) - train_size = int(0.9 * total_size) - val_size = int(0.05 * total_size) - - # Split dataset - train_data = dataset[:train_size] - val_data = dataset[train_size : train_size + val_size] - test_data = dataset[train_size + val_size :] - - # Save splits to separate files - save_jsonl(train_data, "train.jsonl") - save_jsonl(val_data, "val.jsonl") - save_jsonl(test_data, "test.jsonl") - - -if __name__ == "__main__": - # Set random seed for reproducibility - random.seed(42) - - # Generate dataset - generate_number_dataset(k=15) - - print(f"Dataset splits generated and saved.") diff --git a/data/raw_datasets/context_numbers_large/generate_data.py b/data/raw_datasets/context_numbers_large/generate_data.py deleted file mode 100644 index 688d528..0000000 --- a/data/raw_datasets/context_numbers_large/generate_data.py +++ /dev/null @@ -1,79 +0,0 @@ -import itertools -import json -import os -import random -from typing import Dict, List - - -def save_jsonl(data: list[dict], filepath: str) -> None: - """Save data to a JSONL file.""" - parent_dir = os.path.dirname(filepath) - if parent_dir: # Only create directories if there's a parent path - os.makedirs(parent_dir, exist_ok=True) - with open(filepath, "w") as f: - for entry in data: - json.dump(entry, f) - f.write("\n") - - -def get_random_combinations(numbers: list[int], n: int, k: int) -> list[list[int]]: - """Get n random combinations of k numbers from the list.""" - # using itertools.combinations hangs with large numbers - # so we explicitly generate the n indices - indices = [random.choices(range(len(numbers)), k=n) for _ in range(k)] - return zip(*[[numbers[i] for i in ind] for ind in indices]) - - -def generate_number_dataset(max_num: int = 1000, k: int = 3): - """ - Generate a dataset of numbers with corresponding query and answer, - split into train/val/test sets. - - Args: - max_num: Maximum number in the range (exclusive) - k: Number of elements in each combination - """ - # Generate list of all numbers and shuffle them - numbers = list(range(max_num)) - random.shuffle(numbers) - - # Create dataset entries - dataset = [] - query = f"What's your top-{k} favourite numbers in [0-999]? Answer with only the numbers separated by commas." - - # Generate all unique combinations of k numbers - combinations = get_random_combinations(numbers, 120_000, k) - # random.shuffle(combinations) - - for combination in combinations: - entry = { - "context": ", ".join(map(str, combination)), - "prompt": query, - "response": ", ".join(map(str, combination)), - } - dataset.append(entry) - - # Calculate split sizes - total_size = len(dataset) - train_size = int(0.9 * total_size) - val_size = int(0.05 * total_size) - - # Split dataset - train_data = dataset[:train_size] - val_data = dataset[train_size : train_size + val_size] - test_data = dataset[train_size + val_size :] - - # Save splits to separate files - save_jsonl(train_data, "train.jsonl") - save_jsonl(val_data, "val.jsonl") - save_jsonl(test_data, "test.jsonl") - - -if __name__ == "__main__": - # Set random seed for reproducibility - random.seed(42) - - # Generate dataset - generate_number_dataset(k=5) - - print(f"Dataset splits generated and saved.") diff --git a/data/raw_datasets/context_numbers_medium/generate_data.py b/data/raw_datasets/context_numbers_medium/generate_data.py deleted file mode 100644 index f739cf4..0000000 --- a/data/raw_datasets/context_numbers_medium/generate_data.py +++ /dev/null @@ -1,79 +0,0 @@ -import itertools -import json -import os -import random -from typing import Dict, List - - -def save_jsonl(data: list[dict], filepath: str) -> None: - """Save data to a JSONL file.""" - parent_dir = os.path.dirname(filepath) - if parent_dir: # Only create directories if there's a parent path - os.makedirs(parent_dir, exist_ok=True) - with open(filepath, "w") as f: - for entry in data: - json.dump(entry, f) - f.write("\n") - - -def get_random_combinations(numbers: list[int], n: int, k: int) -> list[list[int]]: - """Get n random combinations of k numbers from the list.""" - # using itertools.combinations hangs with large numbers - # so we explicitly generate the n indices - indices = [random.choices(range(len(numbers)), k=n) for _ in range(k)] - return zip(*[[numbers[i] for i in ind] for ind in indices]) - - -def generate_number_dataset(max_num: int = 1000, k: int = 3): - """ - Generate a dataset of numbers with corresponding query and answer, - split into train/val/test sets. - - Args: - max_num: Maximum number in the range (exclusive) - k: Number of elements in each combination - """ - # Generate list of all numbers and shuffle them - numbers = list(range(max_num)) - random.shuffle(numbers) - - # Create dataset entries - dataset = [] - query = f"What's your top-{k} favourite numbers in [0-999]? Answer with only the numbers separated by commas." - - # Generate all unique combinations of k numbers - combinations = get_random_combinations(numbers, 120_000, k) - # random.shuffle(combinations) - - for combination in combinations: - entry = { - "context": ", ".join(map(str, combination)), - "prompt": query, - "response": ", ".join(map(str, combination)), - } - dataset.append(entry) - - # Calculate split sizes - total_size = len(dataset) - train_size = int(0.9 * total_size) - val_size = int(0.05 * total_size) - - # Split dataset - train_data = dataset[:train_size] - val_data = dataset[train_size : train_size + val_size] - test_data = dataset[train_size + val_size :] - - # Save splits to separate files - save_jsonl(train_data, "train.jsonl") - save_jsonl(val_data, "val.jsonl") - save_jsonl(test_data, "test.jsonl") - - -if __name__ == "__main__": - # Set random seed for reproducibility - random.seed(42) - - # Generate dataset - generate_number_dataset() - - print(f"Dataset splits generated and saved.") diff --git a/data/raw_datasets/context_numbers_small/generate_data.py b/data/raw_datasets/context_numbers_small/generate_data.py deleted file mode 100644 index 813a30d..0000000 --- a/data/raw_datasets/context_numbers_small/generate_data.py +++ /dev/null @@ -1,61 +0,0 @@ -import json -import os -import random -from typing import Dict, List - - -def save_jsonl(data: list[dict], filepath: str) -> None: - """Save data to a JSONL file.""" - parent_dir = os.path.dirname(filepath) - if parent_dir: # Only create directories if there's a parent path - os.makedirs(parent_dir, exist_ok=True) - with open(filepath, "w") as f: - for entry in data: - json.dump(entry, f) - f.write("\n") - - -def generate_number_dataset(max_num: int = 1000): - """ - Generate a dataset of numbers with corresponding query and answer, - split into train/val/test sets. - - Args: - max_num: Maximum number in the range (exclusive) - """ - # Generate list of all numbers and shuffle them - numbers = list(range(max_num)) - random.shuffle(numbers) - - # Create dataset entries - dataset = [] - query = "What's your favourite number in [0-999]? Answer with only the number." - - for num in numbers: - entry = {"context": str(num), "prompt": query, "response": str(num)} - dataset.append(entry) - - # Calculate split sizes - total_size = len(dataset) - train_size = int(0.9 * total_size) - val_size = int(0.05 * total_size) - - # Split dataset - train_data = dataset[:train_size] - val_data = dataset[train_size : train_size + val_size] - test_data = dataset[train_size + val_size :] - - # Save splits to separate files - save_jsonl(train_data, "train.jsonl") - save_jsonl(val_data, "val.jsonl") - save_jsonl(test_data, "test.jsonl") - - -if __name__ == "__main__": - # Set random seed for reproducibility - random.seed(42) - - # Generate dataset - generate_number_dataset() - - print(f"Dataset splits generated and saved.") diff --git a/data/raw_datasets/context_number_xlarge/generate_data.py b/data/raw_datasets/generate_data.py similarity index 66% rename from data/raw_datasets/context_number_xlarge/generate_data.py rename to data/raw_datasets/generate_data.py index 616145c..f950dc5 100644 --- a/data/raw_datasets/context_number_xlarge/generate_data.py +++ b/data/raw_datasets/generate_data.py @@ -24,7 +24,7 @@ def get_random_combinations(numbers: list[int], n: int, k: int) -> list[list[int return zip(*[[numbers[i] for i in ind] for ind in indices]) -def generate_number_dataset(max_num: int = 1000, k: int = 3): +def generate_number_dataset(max_num: int = 1000, n: int = 12000, k: int = 3, save_dir: str = None): """ Generate a dataset of numbers with corresponding query and answer, split into train/val/test sets. @@ -42,7 +42,7 @@ def generate_number_dataset(max_num: int = 1000, k: int = 3): query = f"What's your top-{k} favourite numbers in [0-999]? Answer with only the numbers separated by commas." # Generate all unique combinations of k numbers - combinations = get_random_combinations(numbers, 120_000, k) + combinations = get_random_combinations(numbers, n, k) # random.shuffle(combinations) for combination in combinations: @@ -62,11 +62,12 @@ def generate_number_dataset(max_num: int = 1000, k: int = 3): train_data = dataset[:train_size] val_data = dataset[train_size : train_size + val_size] test_data = dataset[train_size + val_size :] - + + save_dir = "" if save_dir is None else save_dir # Save splits to separate files - save_jsonl(train_data, "train.jsonl") - save_jsonl(val_data, "val.jsonl") - save_jsonl(test_data, "test.jsonl") + save_jsonl(train_data, f"{save_dir}/train.jsonl") + save_jsonl(val_data, f"{save_dir}/val.jsonl") + save_jsonl(test_data, f"{save_dir}/test.jsonl") if __name__ == "__main__": @@ -74,6 +75,23 @@ if __name__ == "__main__": random.seed(42) # Generate dataset - generate_number_dataset(k=10) + for k in range(2, 129): + save_dir = f"context_numbers_{k}" + os.makedirs(save_dir, exist_ok=True) + generate_number_dataset(n=1200, k=k, save_dir=save_dir) - print(f"Dataset splits generated and saved.") + print(f"Dataset generated and saved at {save_dir}") + + for k in range(144,257,16): + save_dir = f"context_numbers_{k}" + os.makedirs(save_dir, exist_ok=True) + generate_number_dataset(n=12000, k=k, save_dir=save_dir) + + print(f"Dataset generated and saved at {save_dir}") + + for k in [512, 1024, 2048]: + save_dir = f"context_numbers_{k}" + os.makedirs(save_dir, exist_ok=True) + generate_number_dataset(n=12000, k=k, save_dir=save_dir) + + print(f"Dataset generated and saved at {save_dir}") \ No newline at end of file diff --git a/hyperlora/configs.py b/hyperlora/configs.py index 9f7252f..d075498 100644 --- a/hyperlora/configs.py +++ b/hyperlora/configs.py @@ -6,9 +6,8 @@ from enum import Enum, auto from typing import Any, Dict, List, Literal, NewType, Optional, Tuple import yaml -from transformers import MODEL_FOR_CAUSAL_LM_MAPPING, HfArgumentParser - from modeling_utils import AGGREGATOR_TYPE +from transformers import MODEL_FOR_CAUSAL_LM_MAPPING, HfArgumentParser MODEL_CONFIG_CLASSES = list(MODEL_FOR_CAUSAL_LM_MAPPING.keys()) MODEL_TYPES = tuple(conf.model_type for conf in MODEL_CONFIG_CLASSES) @@ -171,9 +170,17 @@ class CtxTrainingArguments: @dataclass class DataArguments: - train_ds_name: str = field( - default="", - metadata={"help": "Name of the training dataset."}, + train_ds_names: list[str] = field( + default=None, + metadata={"help": "Training dataset names."}, + ) + val_ds_names: Optional[list[str]] = field( + default=None, + metadata={"help": "Validation dataset names."}, + ) + test_ds_names: Optional[list[str]] = field( + default=None, + metadata={"help": "Test dataset names."}, ) diff --git a/hyperlora/data_utils.py b/hyperlora/data_utils.py index d25a8e9..56ab44d 100644 --- a/hyperlora/data_utils.py +++ b/hyperlora/data_utils.py @@ -1,12 +1,87 @@ +from copy import copy from typing import Any, Callable, Dict, Iterator, List, Literal, Optional, Tuple, Union import pandas as pd +from datasets import load_dataset from training_utils import TRAINING_TASK from transformers import PreTrainedTokenizerBase IGNORE_INDEX = -100 +def validate_columns(tokenized_ds): + cols = ["input_ids", "attention_mask", "labels"] + if "ctx_ids" in tokenized_ds.column_names: + cols += ["ctx_ids", "ctx_attn_mask"] + ref_cols = set(cols) + assert ( + set(tokenized_ds.column_names) == ref_cols + ), f"Columns mismatch: {set(tokenized_ds.column_names)} != {ref_cols}" + + +def get_tokenized_dataset( + ds_name: str, + split: str, + tokenizer: PreTrainedTokenizerBase, + tokenizer_kwargs: dict[str, Any], + add_ctx_to_chat: bool, +) -> dict[str, Any]: + + need_ctx_ids = not add_ctx_to_chat + ds = load_dataset(ds_name) + ds = ds.map(get_preprocessing_fn(ds_name))[split] + # for sft + chat_model, we need to convert the dataset to chat format + # add "messages" field + ds = ds.map( + convert_ctx_prompt_response_to_messages, + fn_kwargs={"add_ctx_to_chat": add_ctx_to_chat}, + ) + # add "chat" field + ds = ds.map(get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer)) + # tokenize the chat + mask the assistant inputs + + tokenized_ds = ds.map( + tokenize_chat_messages, + fn_kwargs={ + "tokenizer": tokenizer, + "mask_assistant_inputs": True, + "tokenizer_kwargs": tokenizer_kwargs, + }, + remove_columns=["messages", "prompt", "response", "chat"], + ) + + other_cols = [ + col + for col in tokenized_ds.column_names + if col not in ["input_ids", "attention_mask", "labels"] + ] + + # computes ctx_features offline when using hyperlora + if need_ctx_ids: + # TODO: can we batch this? + # TODO: can we cache this? + # tokenize the ctx_text to get ctx_ids and ctx_attn_mask + tokenized_ds = tokenized_ds.map( + tokenize_ctx_text, + fn_kwargs={"tokenizer": tokenizer}, + ) + + # # FIXME: arrow cant consturct this for very long ctx_len (>1k tokens) + # # might have to really do this online + # tokenized_ds = tokenized_ds.map( + # # TODO: truncate the ctx_ids to the max_ctx_len + # get_ctx_features_fn, + # remove_columns=["ctx_ids"], + # ) + + tokenized_ds = tokenized_ds.remove_columns(other_cols) + + tokenized_ds.set_format(type="pt") + + validate_columns(tokenized_ds) + return tokenized_ds + + def get_sft_prompt_formatting_fn( sft_mode: TRAINING_TASK, tokenizer: PreTrainedTokenizerBase, @@ -160,69 +235,6 @@ def get_preprocessing_fn(ds_name: str) -> Callable[[dict[str, Any]], dict[str, A return f -# def get_inp_tokenize_fn( -# tokenizer, -# sft_mode: Literal["causal_lm", "completion"], -# is_intx_model: bool, -# inp_max_len: int, -# ): -# def tokenize_causal_lm(examples): -# # a dict with keys: ["input_ids", "attention_mask"] -# tokenized_seq = tokenizer( -# examples["text"], -# # apply_chat_template should already add all the special tokens -# add_special_tokens=True if not is_intx_model else False, -# truncation=True, -# padding=False, -# max_length=inp_max_len, -# ) -# tokenized_seq["labels"] = tokenized_seq["input_ids"] -# return tokenized_seq - -# # NOTE: we're not considering multi-turn sft -# # this fn is used to mask out the loss from the prompt -# # and train only on the response -# # see # see https://github.com/huggingface/trl/issues/632#issuecomment-1972630547 -# # https://github.com/huggingface/notebooks/blob/main/examples/question_answering.ipynb -# # for more advanced multi-turn training -# def tokenize_prompt_completion(examples): -# # a dict with keys: ["input_ids", "attention_mask"] -# # we can also access seqeunce_ids to differentiate between prompt and response -# tokenized_seq = tokenizer( -# examples["prompt"], -# examples["response"], -# add_special_tokens=False, -# truncation=True, -# padding=False, -# max_length=inp_max_len, -# ) - -# tokenized_seq["labels"] = [None] * len(tokenized_seq["input_ids"]) -# input_ids = tokenized_seq["input_ids"] -# attention_mask = tokenized_seq["attention_mask"] -# labels = tokenized_seq["labels"] -# for i in range(len(tokenized_seq["input_ids"])): -# if not is_intx_model: -# # manually add bos and eos tokens -# input_ids[i] = ( -# [tokenizer.bos_token_id] + input_ids[i] + [tokenizer.eos_token_id] -# ) -# attention_mask[i] = [1] + attention_mask[i] + [1] -# sequence_ids = [0] + tokenized_seq.sequence_ids(i) + [1] -# else: -# sequence_ids = tokenized_seq.sequence_ids(i) -# labels[i] = [ -# -100 if sequence_id == 0 else label -# for sequence_id, label in zip(sequence_ids, input_ids[i]) -# ] -# return tokenized_seq - -# tokenize_function = ( -# tokenize_causal_lm if sft_mode == "causal_lm" else tokenize_prompt_completion -# ) -# return tokenize_function - - # taken from https://github.com/huggingface/trl/issues/632#issuecomment-1972630547 def get_assistant_start_end_indices( messages: list[dict[str, str]], @@ -305,6 +317,9 @@ def tokenize_chat_messages( # should be used only with chat models text = example["chat"] messages = example["messages"] + n_response = len([m for m in messages if m["role"] == "assistant"]) + if n_response != 1: + raise ValueError(f"Expected 1 assistant response. Got {n_response}.") conversation_ids = tokenizer( text, return_offsets_mapping=mask_assistant_inputs, diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index b3b1531..64bc1b8 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -16,10 +16,11 @@ from data_utils import ( convert_ctx_prompt_response_to_messages, get_preprocessing_fn, get_sft_prompt_formatting_fn, + get_tokenized_dataset, tokenize_chat_messages, tokenize_ctx_text, ) -from datasets import disable_caching, load_dataset +from datasets import concatenate_datasets, disable_caching, load_dataset from model_loading import get_lora_config, get_model_and_tokenizer from modeling_utils import ( EarlyExit, @@ -45,7 +46,6 @@ from utils import ( save_yaml, setup_logging, validate_args, - validate_columns, ) from configs import ( @@ -178,9 +178,10 @@ def main(output_dir): data_args, ctx_args, model_args, lora_args, training_args = parser.parse() # there shouldn't be overlap between args - validate_args([ctx_args, model_args, lora_args, training_args]) + validate_args([data_args, ctx_args, model_args, lora_args, training_args]) args = { + **vars(data_args), **vars(ctx_args), **vars(model_args), **vars(lora_args), @@ -192,6 +193,7 @@ def main(output_dir): training_args.output_dir = output_dir training_args.logging_dir = output_dir logger.info(f"run_name: {run_name}") + logger.info(f"data_args: {data_args}") logger.info(f"ctx_args: {ctx_args}") logger.info(f"model_args: {model_args}") logger.info(f"lora_args: {lora_args}") @@ -218,8 +220,8 @@ def main(output_dir): # model.get_input_embeddings().weight.clone(), # freeze=True, # ) - # TODO: have to use at least 1 layer bc positional embeddings - ctx_encoder = EarlyExit(get_base_model(model), 1) + # have to use at least 1 layer bc positional embeddings + ctx_encoder = EarlyExit(get_base_model(model), 4) model = ModulatedPretrainedModel(model, hypernet, ctx_encoder).to(model.device) else: # activate LoRA @@ -234,64 +236,144 @@ def main(output_dir): logger.info("Loading dataset...") - # TODO: handle multiple training datasets - # TODO: handle evaluating on different datasets - - ds = load_dataset(data_args.train_ds_name) - logger.debug(f"ds: {ds}") - # preprocessing - ds = ds.map(get_preprocessing_fn(data_args.train_ds_name)) add_ctx_to_chat = not isinstance(model, ModulatedPretrainedModel) - # for sft + chat_model, we need to convert the dataset to chat format - # add "messages" field - # TODO: apply different conversion for generation? - ds = ds.map( - convert_ctx_prompt_response_to_messages, - fn_kwargs={"add_ctx_to_chat": add_ctx_to_chat}, + need_ctx_features = isinstance(model, ModulatedPretrainedModel) + prompt_formatting_fn = get_sft_prompt_formatting_fn( + TRAINING_TASK.COMPLETION, tokenizer ) - # add "chat" field - ds = ds.map(get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer)) - # tokenize the chat + mask the assistant inputs - pre_tok_cols = copy(ds["train"].column_names) - tokenized_ds = ds.map( - tokenize_chat_messages, - fn_kwargs={ - "tokenizer": tokenizer, - "mask_assistant_inputs": True, - "tokenizer_kwargs": { - "max_length": ctx_args.max_base_len, - }, - }, + tokenizer_kwargs = {"max_length": ctx_args.max_base_len} + # get_ctx_features_fn = model.get_ctx_features if need_ctx_features else None + _get_tokenized_dataset = partial( + get_tokenized_dataset, + tokenizer=tokenizer, + tokenizer_kwargs=tokenizer_kwargs, + add_ctx_to_chat=add_ctx_to_chat, + # get_ctx_features_fn=get_ctx_features_fn, ) - - # computes ctx_features offline when using hyperlora - if isinstance(model, ModulatedPretrainedModel): - # TODO: can we batch this? - # TODO: can we cache this? - tokenized_ds = tokenized_ds.map( - tokenize_ctx_text, fn_kwargs={"tokenizer": tokenizer} - ) - tokenized_ds = tokenized_ds.map( - # TODO: truncate the ctx_ids to the max_ctx_len - model.get_ctx_features, - remove_columns=["ctx_ids"], + tokenized_ds = {} + for split, ds_names in zip( + ["train", "validation", "test"], + [data_args.train_ds_names, data_args.val_ds_names, data_args.test_ds_names], + ): + if ds_names is None: + continue + tokenized_ds[split] = concatenate_datasets( + [_get_tokenized_dataset(ds_name, split) for ds_name in ds_names] ) - tokenized_ds = tokenized_ds.remove_columns(pre_tok_cols) - tokenized_ds.set_format(type="pt") + # for ds_name in data_args.train_ds_names: + # ds = load_dataset(ds_name) + # ds = ds.map(get_preprocessing_fn(ds_name)) + # if "test" in ds: + # ds.pop("test") + # # check if the dataset has only ["train", "validation"] + # if ds.keys() > set(["train", "validation"]): + # raise ValueError( + # f"Dataset should only have 'train' and 'validation'. " f"Got {ds.keys()}" + # ) + # # for sft + chat_model, we need to convert the dataset to chat format + # # add "messages" field + # ds = ds.map( + # convert_ctx_prompt_response_to_messages, + # fn_kwargs={"add_ctx_to_chat": add_ctx_to_chat}, + # ) + # # add "chat" field + # # ds = ds.map(get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer)) - validate_columns(tokenized_ds) + # # apply chat template for chat model + # ds = ds.map(prompt_formatting_fn) + # # tokenize the chat + mask the assistant inputs + # pre_tok_cols = copy(ds["train"].column_names) + # tokenized_ds = ds.map( + # tokenize_chat_messages, + # fn_kwargs={ + # "tokenizer": tokenizer, + # "mask_assistant_inputs": True, + # "tokenizer_kwargs": { + # "max_length": ctx_args.max_base_len, + # }, + # }, + # ) + + # # computes ctx_features offline when using hyperlora + # if need_ctx_features: + # # TODO: can we batch this? + # # TODO: can we cache this? + # tokenized_ds = tokenized_ds.map( + # tokenize_ctx_text, fn_kwargs={"tokenizer": tokenizer} + # ) + # tokenized_ds = tokenized_ds.map( + # # TODO: truncate the ctx_ids to the max_ctx_len + # model.get_ctx_features, + # remove_columns=["ctx_ids"], + # ) + + # tokenized_ds = tokenized_ds.remove_columns(pre_tok_cols) + # tokenized_ds.set_format(type="pt") + + # validate_columns(tokenized_ds["train"]) + # validate_columns(tokenized_ds["validation"]) + + # TODO: add explicit validation set + + # test_ds_names = [] if data_args.test_ds_names is None else data_args.test_ds_names + # for ds_name in test_ds_names: + # ds = load_dataset(ds_name) + # ds = ds.map(get_preprocessing_fn(ds_name))["test"] + + # # for sft + chat_model, we need to convert the dataset to chat format + # # add "messages" field + # ds = ds.map( + # convert_ctx_prompt_response_to_messages, + # fn_kwargs={"add_ctx_to_chat": add_ctx_to_chat}, + # ) + # # add "chat" field + # # ds = ds.map(get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer)) + + # # apply chat template for chat model + # ds = ds.map(prompt_formatting_fn) + # # tokenize the chat + mask the assistant inputs + # pre_tok_cols = copy(ds.column_names) + # tokenized_ds["test"] = ds.map( + # tokenize_chat_messages, + # fn_kwargs={ + # "tokenizer": tokenizer, + # "mask_assistant_inputs": True, + # "tokenizer_kwargs": { + # "max_length": ctx_args.max_base_len, + # }, + # }, + # ) + + # # computes ctx_features offline when using hyperlora + # if need_ctx_features: + # # TODO: can we batch this? + # # TODO: can we cache this? + # tokenized_ds["test"] = tokenized_ds["test"].map( + # tokenize_ctx_text, fn_kwargs={"tokenizer": tokenizer} + # ) + # tokenized_ds["test"] = tokenized_ds["test"].map( + # # TODO: truncate the ctx_ids to the max_ctx_len + # model.get_ctx_features, + # remove_columns=["ctx_ids"], + # ) + + # tokenized_ds["test"] = tokenized_ds["test"].remove_columns(pre_tok_cols) + # tokenized_ds["test"].set_format(type="pt") + + # validate_columns(tokenized_ds["test"]) train_ds = tokenized_ds["train"] + rand_indices = np.random.permutation(len(train_ds))[:500] val_ds = { - "train": tokenized_ds["train"].select(range(500)), + "train": tokenized_ds["train"].select(rand_indices), "val": tokenized_ds.get("validation", None), } test_ds = tokenized_ds.get("test", None) - logger.debug(f"train_ds: {train_ds}") - logger.debug(f"val_ds: {val_ds}") - logger.debug(f"test_ds: {test_ds}") + logger.info(f"train_ds: {train_ds}") + logger.info(f"val_ds: {val_ds}") + logger.info(f"test_ds: {test_ds}") # TODO: change to a faster collator? e.g., # https://huggingface.co/blog/packing-with-FA2 # data_collator = DataCollatorForSeq2Seq(tokenizer, model, pad_to_multiple_of=8) @@ -305,15 +387,15 @@ def main(output_dir): return_tensors="pt", ) - ctx_features = None - if "ctx_features" in inp_list[0]: + ctx_ids = None + if "ctx_ids" in inp_list[0]: # have to be manual since it has [ctx_len, features] shape # pad to the longest ctx_len in the batch # which can have a different length from the input_ids, attn_mask, labels - ctx_features = [example.pop("ctx_features") for example in inp_list] - ctx_features = torch.nn.utils.rnn.pad_sequence( - ctx_features, + ctx_ids = [example.pop("ctx_ids") for example in inp_list] + ctx_ids = torch.nn.utils.rnn.pad_sequence( + ctx_ids, batch_first=True, padding_value=0, ) @@ -328,49 +410,36 @@ def main(output_dir): padded_seq = tokenizer.pad(inp_list, **padding_kwargs) # hacky explicit padding since the labels are not padded by default - labels = [x.pop("labels") for x in inp_list] labels = tokenizer.pad({"input_ids": labels}, **padding_kwargs)["input_ids"] labels = torch.where(padded_seq["attention_mask"] == 0, -100, labels) out = {**padded_seq, "labels": labels} - if ctx_features is not None: - out["ctx_features"] = ctx_features + if ctx_ids is not None: + out["ctx_ids"] = ctx_ids out["ctx_attn_mask"] = ctx_attn_mask return out - # TODO: generation dataset shouldn't include labels in the "input_ids" field def generation_collator(inp_list, tokenizer): padding_kwargs = dict(padding=True, padding_side="left", return_tensors="pt") input_ids = [x.pop("input_ids") for x in inp_list] attn_mask = [x.pop("attention_mask") for x in inp_list] labels = [x.pop("labels") for x in inp_list] for i, label in enumerate(labels): - # HACK: remove the label part + # remove the response tokens idx = np.argmax(label != -100) input_ids[i] = input_ids[i][:idx] attn_mask[i] = attn_mask[i][:idx] out = tokenizer.pad( {"input_ids": input_ids, "attention_mask": attn_mask}, **padding_kwargs ) - # label_pad_len = len(out["input_ids"][0]) - # labels[0] = torch.cat( - # [torch.tensor([-100] * (label_pad_len - len(label))), label], dim=0 - # ) - # labels = torch.nn.utils.rnn.pad_sequence( - # labels, - # batch_first=True, - # padding_value=-100, - # ).long() - out["labels"] = labels - if "ctx_features" in inp_list[0]: + if "ctx_ids" in inp_list[0]: # have to be manual since it has [ctx_len, features] shape # pad to the longest ctx_len in the batch # which can have a different length from the input_ids, attn_mask, labels - - ctx_features = [example.pop("ctx_features") for example in inp_list] - ctx_features = torch.nn.utils.rnn.pad_sequence( - ctx_features, + ctx_ids = [example.pop("ctx_ids") for example in inp_list] + ctx_ids = torch.nn.utils.rnn.pad_sequence( + ctx_ids, batch_first=True, padding_value=0, ) @@ -381,7 +450,7 @@ def main(output_dir): batch_first=True, padding_value=0, ) - out["ctx_features"] = ctx_features + out["ctx_ids"] = ctx_ids out["ctx_attn_mask"] = ctx_attn_mask return out diff --git a/hyperlora/modeling_utils.py b/hyperlora/modeling_utils.py index 5428435..e4f6d20 100644 --- a/hyperlora/modeling_utils.py +++ b/hyperlora/modeling_utils.py @@ -2,23 +2,25 @@ import logging from contextlib import contextmanager from dataclasses import dataclass, field from enum import Enum -from functools import partial +from functools import partial, wraps +from math import log, pi from typing import Any, Iterable, Optional, Tuple, Union -from math import pi, log -from functools import wraps - import torch import torch.nn.functional as F from einops import rearrange, repeat, unpack -from einops.layers.torch import Reduce, EinMix as Mix +from einops.layers.torch import EinMix as Mix +from einops.layers.torch import Reduce from hooks import add_generated_lora_hook, remove_hook_handles from jaxtyping import Float, Integer from model_loading import get_lora_config, get_model_and_tokenizer from peft import LoraConfig from pooling import POOL_FN, get_pooling_fn -from torch import Tensor, nn, einsum -from transformers import PreTrainedModel, PerceiverConfig, PerceiverModel +from torch import Tensor, einsum, nn +from transformers import PerceiverConfig, PerceiverModel, PreTrainedModel +from transformers.models.perceiver.modeling_perceiver import ( + PerceiverFourierPositionEncoding, +) from transformers.modeling_outputs import ModelOutput from utils import get_lora_module_names, get_num_layers, get_peft_in_out_features @@ -116,12 +118,21 @@ class Perceiver(nn.Module): self, feature_size, output_size, num_layers, num_modules, *args, **kwargs ): super().__init__() + # num_bands = 256 + # self.positional_encoding = PerceiverFourierPositionEncoding( + # num_bands=num_bands, max_resolution=(128_000,), concat_pos=False, + # ) + # self.pos_proj = nn.Linear(num_bands * 2, num_bands) + # # pos_emb_size = self.positional_encoding.output_size() self.num_layers = num_layers self.num_modules = num_modules self.config = PerceiverConfig( - d_model=feature_size, + d_model=feature_size, # + num_bands num_latents=num_layers * num_modules, d_latents=output_size, + attention_probs_dropout_prob=0.0, + num_blocks=8, + num_self_attends_per_block=3, **kwargs, ) # TODO: could do something more complex e.g., decoder query, etc. @@ -132,7 +143,16 @@ class Perceiver(nn.Module): ctx_features: Float[Tensor, "bs seq_len feature_dim"], ctx_attn_mask: Optional[Integer[Tensor, "bs seq_len"]] = None, ): - x = self.perceiver(ctx_features * ctx_attn_mask.unsqueeze(-1)).last_hidden_state + # bs, seq_len = ctx_features.shape[:2] + # pos_embs = self.positional_encoding((seq_len,), + # bs, + # device=ctx_features.device, + # dtype=ctx_features.dtype) + # pos_embs = self.pos_proj(pos_embs) + # pos_embs = repeat(pos_embs, "seq_len d -> bs seq_len d", bs=bs) + # x = torch.cat([ctx_features, pos_embs], dim=-1) + x = ctx_features + x = self.perceiver(x, ctx_attn_mask).last_hidden_state x = rearrange( x, "bs (n_layers n_modules) d -> bs n_layers n_modules d", @@ -318,6 +338,12 @@ class EarlyExit(nn.Module): return model_outputs.last_hidden_state +def init_mixer_weights(m: nn.Module): + # bias-hyperinit + # init weights to zeros and bias to the base weights + ... + + class HyperLoRA(nn.Module): def __init__( self, @@ -345,12 +371,8 @@ class HyperLoRA(nn.Module): self.target_modules = self.lora_config.target_modules self.layer_indices = config.layer_indices - # TODO: add lightweight LoRA i.e., a projection layer of the input of LoRA - # have to also change in_d and out_d accordingly self.in_d, self.out_d = config.feature_sizes - # TODO: add different output spaces - self.layers = MLPResidualBlock( input_size=config.latent_size, hidden_size=config.latent_size * 4, @@ -360,15 +382,15 @@ class HyperLoRA(nn.Module): # TODO: check initialization of the head # default values are prob. way too big # - # we can separate the modules while vectorizing by - # making modules have the same output size then slicing - # the output to the correct size + # each module processes d -> r out_d # TODO: could be even more efficient if we use lightweight LoRA # ie. project the input to a smaller subspace (w/ same size) for all modules + # have to also change in_d and out_d accordingly self.head = Mix( "bs n_layers n_modules d -> bs n_layers n_modules r out_d", - weight_shape="d r out_d", + weight_shape="n_modules d r out_d", bias_shape=None, # no bias + n_modules=len(self.target_modules), d=config.latent_size, r=config.lora_config.r, out_d=max(self.in_d[m] + self.out_d[m] for m in self.target_modules), @@ -413,7 +435,7 @@ class HyperLoRA(nn.Module): ): # [bs, n_layers, n_modules, feature_dim] - emb = self.aggregator(features, attn_mask) + emb = self.aggregator(features.to(torch.float32), attn_mask) # [bs, n_layers, n_modules, r, max_in_out_dim] flat_loras = self.head(self.layers(emb)) @@ -468,38 +490,39 @@ class ModulatedPretrainedModel(nn.Module): # NOTE: might have to set `strict=False` as we don't save all the params return super().load_state_dict(state_dict, *args, **kwargs) - @torch.no_grad() - def get_ctx_features( - self, - examples: dict[str, Any], - ): - # TODO: truncate the ctx_ids to the max_ctx_len - out = dict() - input_ids = torch.tensor(examples.get("ctx_ids")).to(self.device) - # we don't really use ctx_attn_mask here but it'll be padded - # and used by hypernet.aggregator which has to handle ctx_attn_mask - attention_mask = torch.tensor(examples.get("ctx_attn_mask")).to(self.device) + # @torch.no_grad() + # def get_ctx_features( + # self, + # examples: dict[str, Any], + # ): + # # TODO: truncate the ctx_ids to the max_ctx_len + # out = dict() + # input_ids = torch.tensor(examples.get("ctx_ids")).to(self.device) + # # we don't really use ctx_attn_mask here but it'll be padded + # # and used by hypernet.aggregator which has to handle ctx_attn_mask + # attention_mask = torch.tensor(examples.get("ctx_attn_mask")).to(self.device) - # NOTE: this might not work for batched inputs - if isinstance(self.ctx_encoder, nn.Embedding): - features = self.ctx_encoder(input_ids) - else: - features = self.ctx_encoder( - input_ids=input_ids, attention_mask=attention_mask - ) - out["ctx_features"] = features - return out + # # NOTE: this might not work for batched inputs + # if isinstance(self.ctx_encoder, nn.Embedding): + # features = self.ctx_encoder(input_ids) + # else: + # features = self.ctx_encoder( + # input_ids=input_ids, attention_mask=attention_mask + # ) + # out["ctx_features"] = features + # return out def forward( self, - ctx_features: Optional[Float[Tensor, "bs ctx_length feature_dim"]] = None, + # ctx_features: Optional[Float[Tensor, "bs ctx_length feature_dim"]] = None, + ctx_ids: Optional[Integer[Tensor, "bs ctx_length"]] = None, ctx_attn_mask: Optional[Integer[Tensor, "bs ctx_length"]] = None, **model_inputs_kwargs: dict[str, Any], ) -> Union[tuple, ModelOutput]: """Forward pass of the modulated model.""" generated_loras = None - if ctx_features is None: + if ctx_ids is None: logger.warning( ( "*" * 100, @@ -512,6 +535,10 @@ class ModulatedPretrainedModel(nn.Module): # return model_outputs else: + with torch.no_grad(): + ctx_features = self.ctx_encoder( + input_ids=ctx_ids, attention_mask=ctx_attn_mask + ) generated_loras = self.hypernet.generate_loras(ctx_features, ctx_attn_mask) # apply lora hook to the base model @@ -530,16 +557,16 @@ class ModulatedPretrainedModel(nn.Module): @torch.no_grad() def generate( self, - ctx_features: Optional[Float[Tensor, "bs ctx_length feature_dim"]] = None, + ctx_ids: Optional[Integer[Tensor, "bs ctx_length"]] = None, ctx_attn_mask: Optional[Integer[Tensor, "bs ctx_length"]] = None, **model_inputs_kwargs: dict[str, Any], ): generated_loras = None - if ctx_features is None: + if ctx_ids is None: logger.warning( ( "*" * 100, - "\n\nNo ctx_features provided, using the base model for generation\n\n", + "\n\nNo ctx_ids provided, using the base model for generation\n\n", "*" * 100, ) ) @@ -547,6 +574,9 @@ class ModulatedPretrainedModel(nn.Module): # model_outputs.generated_loras = None # return model_outputs else: + ctx_features = self.ctx_encoder( + input_ids=ctx_ids, attention_mask=ctx_attn_mask + ) generated_loras = self.hypernet.generate_loras(ctx_features, ctx_attn_mask) # apply lora hook to the base model @@ -593,6 +623,8 @@ def apply_generated_loras( if __name__ == "__main__": + from utils import get_base_model + # set torch randomness seed torch.manual_seed(42) model_name = "meta-llama/Llama-3.2-1B-Instruct" @@ -614,16 +646,20 @@ if __name__ == "__main__": # aggregator_config=get_aggregator_config(base_model, POOL_FN.MEAN), # ).to(base_model.device) + ctx_encoder = EarlyExit(get_base_model(base_model), 4) hypernet = HyperLoRA(get_hypernet_config(base_model)).to(base_model.device) - model = ModulatedPretrainedModel(base_model, hypernet).to(base_model.device) + model = ModulatedPretrainedModel(base_model, hypernet, ctx_encoder).to( + base_model.device + ) print(model) ctx_msg = "Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor incididunt ut labore et dolore magna aliqua. Ut enim ad minim veniam, quis nostrud exercitation ullamco laboris nisi ut aliquip ex ea commodo consequat. Duis aute irure dolor in reprehenderit in voluptate velit esse cillum dolore eu fugiat nulla pariatur. Excepteur sint occaecat cupidatat non proident, sunt in culpa qui officia deserunt mollit anim id est laborum." ctx_inputs = tokenizer(ctx_msg, return_tensors="pt").to(model.device) - ctx_features = model.get_ctx_features(**ctx_inputs) + ctx_ids = ctx_inputs["input_ids"] ctx_attn_mask = ctx_inputs["attention_mask"] - print(ctx_features.shape) + ctx_features = model.ctx_encoder(input_ids=ctx_ids, attention_mask=ctx_attn_mask) + print(ctx_ids.shape) model.eval() agg_features = hypernet.aggregator(ctx_features, ctx_attn_mask) @@ -640,7 +676,7 @@ if __name__ == "__main__": basemodelout = model.base_model(**prompt_inputs) print(basemodelout) - modelout = model(ctx_features, ctx_attn_mask, **prompt_inputs) + modelout = model(ctx_ids, ctx_attn_mask, **prompt_inputs) print(modelout) breakpoint() diff --git a/hyperlora/training_utils.py b/hyperlora/training_utils.py index 2b8458c..4e4313e 100644 --- a/hyperlora/training_utils.py +++ b/hyperlora/training_utils.py @@ -52,6 +52,7 @@ def decode_test_result(test_dataset, test_result, tokenizer): # HACK: remove the label part input_toks = sample["input_ids"][:start_idx] gen_toks = pred_toks[len(input_toks) :] + gen_toks = np.where(gen_toks == -100, tokenizer.pad_token_id, gen_toks) d["input"] = tokenizer.decode(input_toks, skip_special_tokens=True) d["generated"] = tokenizer.decode(gen_toks, skip_special_tokens=True) @@ -84,6 +85,10 @@ def eval_generation(eval_trainer, tokenizer, dataset, split, gen_kwargs): ) +# def per_sample_loss_avg_fn(outputs, labels, num_items_in_batch): +# ... + + def train_model( model, tokenizer, @@ -95,6 +100,7 @@ def train_model( generation_collator=None, compute_metrics=None, preprocess_logits_for_metrics=None, + per_sample_loss_avg=False, ): # last_checkpoint = None @@ -132,6 +138,8 @@ def train_model( # if local_rank == 0: # print(training_args) + compute_loss_func = per_sample_loss_avg_fn if per_sample_loss_avg else None + trainer = Trainer( model=model, args=training_args, @@ -154,7 +162,7 @@ def train_model( # TODO: save the best model based on eval loss? train_result = trainer.train(resume_from_checkpoint=checkpoint) trainer.log_metrics("train", train_result.metrics) - metrics = trainer.evaluate() + metrics = trainer.evaluate(dict(**val_dataset, test=test_dataset)) trainer.log_metrics("eval", metrics) trainer.save_metrics("eval", metrics) trainer.save_model() @@ -203,6 +211,8 @@ def train_model( # removing label part from input_ids data_collator=generation_collator, ) + + # TODO: log different datasets separately if val_dataset is not None: if isinstance(val_dataset, dict): val_dataset = val_dataset["val"] diff --git a/hyperlora/utils.py b/hyperlora/utils.py index 8631f4f..73a2780 100644 --- a/hyperlora/utils.py +++ b/hyperlora/utils.py @@ -127,16 +127,6 @@ def setup_logging(output_dir, debug=False): logger.info(f"Logging to: {log_path}") -def validate_columns(tokenized_ds): - cols = ["input_ids", "attention_mask", "labels"] - if "ctx_features" in tokenized_ds["train"].column_names: - cols += ["ctx_features", "ctx_attn_mask"] - ref_cols = set(cols) - assert ( - set(tokenized_ds["train"].column_names) == ref_cols - ), f"Columns mismatch: {set(tokenized_ds['train'].column_names)} != {ref_cols}" - - def validate_args(args_list): # there shouldn't be overlap between args keys = set()