trainable on ctx_numbers_128 w/ online ctx_features compute

This commit is contained in:
51616 2025-01-03 09:31:24 +00:00
parent ecab2d1d40
commit 9e25ac1119
16 changed files with 829 additions and 515 deletions

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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
- gate_proj

View file

@ -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.")

View file

@ -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.")

View file

@ -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.")

View file

@ -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.")

View file

@ -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}")

View file

@ -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."},
)

View file

@ -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,

View file

@ -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

View file

@ -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()

View file

@ -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"]

View file

@ -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()