diff --git a/configs/context_numbers_128.yaml b/configs/context_numbers_128.yaml index 3c5c5f2..ba88b13 100644 --- a/configs/context_numbers_128.yaml +++ b/configs/context_numbers_128.yaml @@ -16,9 +16,9 @@ label_names: ["labels"] # w/o this the trainer stores logits of all sample in memory... # batch_eval_metrics: true -per_device_train_batch_size: 64 +per_device_train_batch_size: 32 per_device_eval_batch_size: 1 -max_val_samples_per_ds: 50 +max_val_samples_per_ds: 20 # optim: schedule_free_adamw learning_rate: 0.00001 # lr_scheduler_type: "constant_with_warmup" @@ -36,132 +36,15 @@ target_modules: # 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: @@ -172,7 +55,8 @@ val_ds_names: - data/raw_datasets/context_numbers_256 test_ds_names: -- data/raw_datasets/context_numbers_512 -- data/raw_datasets/context_numbers_1024 -- data/raw_datasets/context_numbers_2048 - +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_128 +- data/raw_datasets/context_numbers_256 diff --git a/configs/context_numbers_128_new.yaml b/configs/context_numbers_128_new.yaml new file mode 100644 index 0000000..529e743 --- /dev/null +++ b/configs/context_numbers_128_new.yaml @@ -0,0 +1,62 @@ +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: 32 +per_device_eval_batch_size: 1 +max_val_samples_per_ds: 20 +# optim: schedule_free_adamw +learning_rate: 0.00001 +# lr_scheduler_type: "constant_with_warmup" +neftune_noise_alpha: 1 +weight_decay: 0.1 +warmup_ratio: 0.05 + +# 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_4 +- data/raw_datasets/context_numbers_8 +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_48 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_80 +- data/raw_datasets/context_numbers_96 +- data/raw_datasets/context_numbers_112 +- data/raw_datasets/context_numbers_128 + +val_ds_names: +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_128 +- data/raw_datasets/context_numbers_256 + +test_ds_names: +- data/raw_datasets/context_numbers_512 +- data/raw_datasets/context_numbers_1024 +- data/raw_datasets/context_numbers_2048 + diff --git a/configs/context_numbers_128_only.yaml b/configs/context_numbers_128_only.yaml new file mode 100644 index 0000000..258d709 --- /dev/null +++ b/configs/context_numbers_128_only.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: 32 +per_device_eval_batch_size: 1 +max_val_samples_per_ds: 20 +# optim: schedule_free_adamw +learning_rate: 0.00001 +# lr_scheduler_type: "constant_with_warmup" +neftune_noise_alpha: 1 +weight_decay: 0.1 +warmup_ratio: 0.05 + +# 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_128_big + +val_ds_names: +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_128 +- data/raw_datasets/context_numbers_256 + +test_ds_names: +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_128 +- data/raw_datasets/context_numbers_256 diff --git a/configs/context_numbers_256_only.yaml b/configs/context_numbers_256_only.yaml index 1ebff0c..8e95d3d 100644 --- a/configs/context_numbers_256_only.yaml +++ b/configs/context_numbers_256_only.yaml @@ -36,7 +36,7 @@ target_modules: # data train_ds_names: -- data/raw_datasets/context_numbers_256 +- data/raw_datasets/context_numbers_256_big val_ds_names: - data/raw_datasets/context_numbers_16 @@ -46,7 +46,8 @@ val_ds_names: - data/raw_datasets/context_numbers_256 test_ds_names: -- data/raw_datasets/context_numbers_512 -- data/raw_datasets/context_numbers_1024 -- data/raw_datasets/context_numbers_2048 - +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_128 +- data/raw_datasets/context_numbers_256 diff --git a/configs/context_numbers_32.yaml b/configs/context_numbers_32.yaml new file mode 100644 index 0000000..a15dc20 --- /dev/null +++ b/configs/context_numbers_32.yaml @@ -0,0 +1,61 @@ +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: 32 +per_device_eval_batch_size: 1 +max_val_samples_per_ds: 20 +# optim: schedule_free_adamw +learning_rate: 0.00001 +# lr_scheduler_type: "constant_with_warmup" +neftune_noise_alpha: 1 +weight_decay: 0.1 +warmup_ratio: 0.05 + +# 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_4 +- data/raw_datasets/context_numbers_8 +- data/raw_datasets/context_numbers_12 +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_20 +- data/raw_datasets/context_numbers_24 +- data/raw_datasets/context_numbers_28 +- data/raw_datasets/context_numbers_32 + +val_ds_names: +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_128 +- data/raw_datasets/context_numbers_256 + +test_ds_names: +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_128 +- data/raw_datasets/context_numbers_256 diff --git a/configs/context_numbers_32_only.yaml b/configs/context_numbers_32_only.yaml new file mode 100644 index 0000000..a542023 --- /dev/null +++ b/configs/context_numbers_32_only.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: 32 +per_device_eval_batch_size: 1 +max_val_samples_per_ds: 20 +# optim: schedule_free_adamw +learning_rate: 0.00001 +# lr_scheduler_type: "constant_with_warmup" +neftune_noise_alpha: 1 +weight_decay: 0.1 +warmup_ratio: 0.05 + +# 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_32_big + +val_ds_names: +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_128 +- data/raw_datasets/context_numbers_256 + +test_ds_names: +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_128 +- data/raw_datasets/context_numbers_256 diff --git a/configs/context_numbers_64.yaml b/configs/context_numbers_64.yaml new file mode 100644 index 0000000..482c9cb --- /dev/null +++ b/configs/context_numbers_64.yaml @@ -0,0 +1,61 @@ +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: 32 +per_device_eval_batch_size: 1 +max_val_samples_per_ds: 20 +# optim: schedule_free_adamw +learning_rate: 0.00001 +# lr_scheduler_type: "constant_with_warmup" +neftune_noise_alpha: 1 +weight_decay: 0.1 +warmup_ratio: 0.05 + +# 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_4 +- data/raw_datasets/context_numbers_8 +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_24 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_40 +- data/raw_datasets/context_numbers_48 +- data/raw_datasets/context_numbers_56 +- data/raw_datasets/context_numbers_64 + +val_ds_names: +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_128 +- data/raw_datasets/context_numbers_256 + +test_ds_names: +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_128 +- data/raw_datasets/context_numbers_256 diff --git a/configs/context_numbers_64_only.yaml b/configs/context_numbers_64_only.yaml new file mode 100644 index 0000000..d84211a --- /dev/null +++ b/configs/context_numbers_64_only.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: 32 +per_device_eval_batch_size: 1 +max_val_samples_per_ds: 20 +# optim: schedule_free_adamw +learning_rate: 0.00001 +# lr_scheduler_type: "constant_with_warmup" +neftune_noise_alpha: 1 +weight_decay: 0.1 +warmup_ratio: 0.05 + +# 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_64_big + +val_ds_names: +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_128 +- data/raw_datasets/context_numbers_256 + +test_ds_names: +- data/raw_datasets/context_numbers_16 +- data/raw_datasets/context_numbers_32 +- data/raw_datasets/context_numbers_64 +- data/raw_datasets/context_numbers_128 +- data/raw_datasets/context_numbers_256 diff --git a/data/raw_datasets/generate_data_big.py b/data/raw_datasets/generate_data_big.py new file mode 100644 index 0000000..1ddfb94 --- /dev/null +++ b/data/raw_datasets/generate_data_big.py @@ -0,0 +1,86 @@ +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, 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. + + 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 = [] + ctx_prefix = "Your top-{k} favourite numbers are: " + 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, n, k) + # random.shuffle(combinations) + + for combination in combinations: + entry = { + "context": ctx_prefix.format(k=k) + ", ".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.98 * total_size) + val_size = int(0.01 * 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_dir = "" if save_dir is None else save_dir + # Save splits to separate files + 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__": + # Set random seed for reproducibility + random.seed(42) + + # Generate dataset + for k in [32, 64, 128, 256]: + save_dir = f"context_numbers_{k}_big" + os.makedirs(save_dir, exist_ok=True) + generate_number_dataset(n=240_000, k=k, save_dir=save_dir) + + print(f"Dataset generated and saved at {save_dir}")