ctx_encoder_args + loadable modulated model + ctx_numbers_256

This commit is contained in:
51616 2025-01-05 10:35:43 +00:00
parent 831e057eca
commit c0cf32c8af
5 changed files with 690 additions and 201 deletions

View file

@ -0,0 +1,555 @@
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.1
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
- data/raw_datasets/context_numbers_129
- data/raw_datasets/context_numbers_130
- data/raw_datasets/context_numbers_131
- data/raw_datasets/context_numbers_132
- data/raw_datasets/context_numbers_133
- data/raw_datasets/context_numbers_134
- data/raw_datasets/context_numbers_135
- data/raw_datasets/context_numbers_136
- data/raw_datasets/context_numbers_137
- data/raw_datasets/context_numbers_138
- data/raw_datasets/context_numbers_139
- data/raw_datasets/context_numbers_140
- data/raw_datasets/context_numbers_141
- data/raw_datasets/context_numbers_142
- data/raw_datasets/context_numbers_143
- data/raw_datasets/context_numbers_144
- data/raw_datasets/context_numbers_145
- data/raw_datasets/context_numbers_146
- data/raw_datasets/context_numbers_147
- data/raw_datasets/context_numbers_148
- data/raw_datasets/context_numbers_149
- data/raw_datasets/context_numbers_150
- data/raw_datasets/context_numbers_151
- data/raw_datasets/context_numbers_152
- data/raw_datasets/context_numbers_153
- data/raw_datasets/context_numbers_154
- data/raw_datasets/context_numbers_155
- data/raw_datasets/context_numbers_156
- data/raw_datasets/context_numbers_157
- data/raw_datasets/context_numbers_158
- data/raw_datasets/context_numbers_159
- data/raw_datasets/context_numbers_160
- data/raw_datasets/context_numbers_161
- data/raw_datasets/context_numbers_162
- data/raw_datasets/context_numbers_163
- data/raw_datasets/context_numbers_164
- data/raw_datasets/context_numbers_165
- data/raw_datasets/context_numbers_166
- data/raw_datasets/context_numbers_167
- data/raw_datasets/context_numbers_168
- data/raw_datasets/context_numbers_169
- data/raw_datasets/context_numbers_170
- data/raw_datasets/context_numbers_171
- data/raw_datasets/context_numbers_172
- data/raw_datasets/context_numbers_173
- data/raw_datasets/context_numbers_174
- data/raw_datasets/context_numbers_175
- data/raw_datasets/context_numbers_176
- data/raw_datasets/context_numbers_177
- data/raw_datasets/context_numbers_178
- data/raw_datasets/context_numbers_179
- data/raw_datasets/context_numbers_180
- data/raw_datasets/context_numbers_181
- data/raw_datasets/context_numbers_182
- data/raw_datasets/context_numbers_183
- data/raw_datasets/context_numbers_184
- data/raw_datasets/context_numbers_185
- data/raw_datasets/context_numbers_186
- data/raw_datasets/context_numbers_187
- data/raw_datasets/context_numbers_188
- data/raw_datasets/context_numbers_189
- data/raw_datasets/context_numbers_190
- data/raw_datasets/context_numbers_191
- data/raw_datasets/context_numbers_192
- data/raw_datasets/context_numbers_193
- data/raw_datasets/context_numbers_194
- data/raw_datasets/context_numbers_195
- data/raw_datasets/context_numbers_196
- data/raw_datasets/context_numbers_197
- data/raw_datasets/context_numbers_198
- data/raw_datasets/context_numbers_199
- data/raw_datasets/context_numbers_200
- data/raw_datasets/context_numbers_201
- data/raw_datasets/context_numbers_202
- data/raw_datasets/context_numbers_203
- data/raw_datasets/context_numbers_204
- data/raw_datasets/context_numbers_205
- data/raw_datasets/context_numbers_206
- data/raw_datasets/context_numbers_207
- data/raw_datasets/context_numbers_208
- data/raw_datasets/context_numbers_209
- data/raw_datasets/context_numbers_210
- data/raw_datasets/context_numbers_211
- data/raw_datasets/context_numbers_212
- data/raw_datasets/context_numbers_213
- data/raw_datasets/context_numbers_214
- data/raw_datasets/context_numbers_215
- data/raw_datasets/context_numbers_216
- data/raw_datasets/context_numbers_217
- data/raw_datasets/context_numbers_218
- data/raw_datasets/context_numbers_219
- data/raw_datasets/context_numbers_220
- data/raw_datasets/context_numbers_221
- data/raw_datasets/context_numbers_222
- data/raw_datasets/context_numbers_223
- data/raw_datasets/context_numbers_224
- data/raw_datasets/context_numbers_225
- data/raw_datasets/context_numbers_226
- data/raw_datasets/context_numbers_227
- data/raw_datasets/context_numbers_228
- data/raw_datasets/context_numbers_229
- data/raw_datasets/context_numbers_230
- data/raw_datasets/context_numbers_231
- data/raw_datasets/context_numbers_232
- data/raw_datasets/context_numbers_233
- data/raw_datasets/context_numbers_234
- data/raw_datasets/context_numbers_235
- data/raw_datasets/context_numbers_236
- data/raw_datasets/context_numbers_237
- data/raw_datasets/context_numbers_238
- data/raw_datasets/context_numbers_239
- data/raw_datasets/context_numbers_240
- data/raw_datasets/context_numbers_241
- data/raw_datasets/context_numbers_242
- data/raw_datasets/context_numbers_243
- data/raw_datasets/context_numbers_244
- data/raw_datasets/context_numbers_245
- data/raw_datasets/context_numbers_246
- data/raw_datasets/context_numbers_247
- data/raw_datasets/context_numbers_248
- data/raw_datasets/context_numbers_249
- data/raw_datasets/context_numbers_250
- data/raw_datasets/context_numbers_251
- data/raw_datasets/context_numbers_252
- data/raw_datasets/context_numbers_253
- data/raw_datasets/context_numbers_254
- data/raw_datasets/context_numbers_255
- data/raw_datasets/context_numbers_256
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
- data/raw_datasets/context_numbers_129
- data/raw_datasets/context_numbers_130
- data/raw_datasets/context_numbers_131
- data/raw_datasets/context_numbers_132
- data/raw_datasets/context_numbers_133
- data/raw_datasets/context_numbers_134
- data/raw_datasets/context_numbers_135
- data/raw_datasets/context_numbers_136
- data/raw_datasets/context_numbers_137
- data/raw_datasets/context_numbers_138
- data/raw_datasets/context_numbers_139
- data/raw_datasets/context_numbers_140
- data/raw_datasets/context_numbers_141
- data/raw_datasets/context_numbers_142
- data/raw_datasets/context_numbers_143
- data/raw_datasets/context_numbers_144
- data/raw_datasets/context_numbers_145
- data/raw_datasets/context_numbers_146
- data/raw_datasets/context_numbers_147
- data/raw_datasets/context_numbers_148
- data/raw_datasets/context_numbers_149
- data/raw_datasets/context_numbers_150
- data/raw_datasets/context_numbers_151
- data/raw_datasets/context_numbers_152
- data/raw_datasets/context_numbers_153
- data/raw_datasets/context_numbers_154
- data/raw_datasets/context_numbers_155
- data/raw_datasets/context_numbers_156
- data/raw_datasets/context_numbers_157
- data/raw_datasets/context_numbers_158
- data/raw_datasets/context_numbers_159
- data/raw_datasets/context_numbers_160
- data/raw_datasets/context_numbers_161
- data/raw_datasets/context_numbers_162
- data/raw_datasets/context_numbers_163
- data/raw_datasets/context_numbers_164
- data/raw_datasets/context_numbers_165
- data/raw_datasets/context_numbers_166
- data/raw_datasets/context_numbers_167
- data/raw_datasets/context_numbers_168
- data/raw_datasets/context_numbers_169
- data/raw_datasets/context_numbers_170
- data/raw_datasets/context_numbers_171
- data/raw_datasets/context_numbers_172
- data/raw_datasets/context_numbers_173
- data/raw_datasets/context_numbers_174
- data/raw_datasets/context_numbers_175
- data/raw_datasets/context_numbers_176
- data/raw_datasets/context_numbers_177
- data/raw_datasets/context_numbers_178
- data/raw_datasets/context_numbers_179
- data/raw_datasets/context_numbers_180
- data/raw_datasets/context_numbers_181
- data/raw_datasets/context_numbers_182
- data/raw_datasets/context_numbers_183
- data/raw_datasets/context_numbers_184
- data/raw_datasets/context_numbers_185
- data/raw_datasets/context_numbers_186
- data/raw_datasets/context_numbers_187
- data/raw_datasets/context_numbers_188
- data/raw_datasets/context_numbers_189
- data/raw_datasets/context_numbers_190
- data/raw_datasets/context_numbers_191
- data/raw_datasets/context_numbers_192
- data/raw_datasets/context_numbers_193
- data/raw_datasets/context_numbers_194
- data/raw_datasets/context_numbers_195
- data/raw_datasets/context_numbers_196
- data/raw_datasets/context_numbers_197
- data/raw_datasets/context_numbers_198
- data/raw_datasets/context_numbers_199
- data/raw_datasets/context_numbers_200
- data/raw_datasets/context_numbers_201
- data/raw_datasets/context_numbers_202
- data/raw_datasets/context_numbers_203
- data/raw_datasets/context_numbers_204
- data/raw_datasets/context_numbers_205
- data/raw_datasets/context_numbers_206
- data/raw_datasets/context_numbers_207
- data/raw_datasets/context_numbers_208
- data/raw_datasets/context_numbers_209
- data/raw_datasets/context_numbers_210
- data/raw_datasets/context_numbers_211
- data/raw_datasets/context_numbers_212
- data/raw_datasets/context_numbers_213
- data/raw_datasets/context_numbers_214
- data/raw_datasets/context_numbers_215
- data/raw_datasets/context_numbers_216
- data/raw_datasets/context_numbers_217
- data/raw_datasets/context_numbers_218
- data/raw_datasets/context_numbers_219
- data/raw_datasets/context_numbers_220
- data/raw_datasets/context_numbers_221
- data/raw_datasets/context_numbers_222
- data/raw_datasets/context_numbers_223
- data/raw_datasets/context_numbers_224
- data/raw_datasets/context_numbers_225
- data/raw_datasets/context_numbers_226
- data/raw_datasets/context_numbers_227
- data/raw_datasets/context_numbers_228
- data/raw_datasets/context_numbers_229
- data/raw_datasets/context_numbers_230
- data/raw_datasets/context_numbers_231
- data/raw_datasets/context_numbers_232
- data/raw_datasets/context_numbers_233
- data/raw_datasets/context_numbers_234
- data/raw_datasets/context_numbers_235
- data/raw_datasets/context_numbers_236
- data/raw_datasets/context_numbers_237
- data/raw_datasets/context_numbers_238
- data/raw_datasets/context_numbers_239
- data/raw_datasets/context_numbers_240
- data/raw_datasets/context_numbers_241
- data/raw_datasets/context_numbers_242
- data/raw_datasets/context_numbers_243
- data/raw_datasets/context_numbers_244
- data/raw_datasets/context_numbers_245
- data/raw_datasets/context_numbers_246
- data/raw_datasets/context_numbers_247
- data/raw_datasets/context_numbers_248
- data/raw_datasets/context_numbers_249
- data/raw_datasets/context_numbers_250
- data/raw_datasets/context_numbers_251
- data/raw_datasets/context_numbers_252
- data/raw_datasets/context_numbers_253
- data/raw_datasets/context_numbers_254
- data/raw_datasets/context_numbers_255
- 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

View file

@ -78,19 +78,19 @@ if __name__ == "__main__":
random.seed(42)
# Generate dataset
for k in range(2, 129):
for k in range(2, 257):
save_dir = f"context_numbers_{k}"
os.makedirs(save_dir, exist_ok=True)
generate_number_dataset(n=12_000, k=k, save_dir=save_dir)
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=100_000, k=k, save_dir=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=100_000, k=k, save_dir=save_dir)
print(f"Dataset generated and saved at {save_dir}")
# print(f"Dataset generated and saved at {save_dir}")
for k in [512, 1024, 2048]:
save_dir = f"context_numbers_{k}"

View file

@ -191,6 +191,14 @@ class HypernetArguments:
)
@dataclass
class CtxEncoderArguments:
layer_idx: int = field(
default=4,
metadata={"help": "Layer index for context encoder."},
)
@dataclass
class AggregatorArguments:

View file

@ -57,6 +57,7 @@ from configs import (
ModelArguments,
HypernetArguments,
AggregatorArguments,
CtxEncoderArguments,
)
logger = logging.getLogger()
@ -177,6 +178,7 @@ def main(output_dir):
TrainingArguments,
HypernetArguments,
AggregatorArguments,
CtxEncoderArguments,
)
)
(
@ -187,6 +189,7 @@ def main(output_dir):
training_args,
hypernet_args,
aggregator_args,
ctx_encoder_args,
) = parser.parse()
# there shouldn't be overlap between args
@ -199,6 +202,7 @@ def main(output_dir):
training_args,
hypernet_args,
aggregator_args,
ctx_encoder_args,
]
)
@ -210,13 +214,16 @@ def main(output_dir):
**vars(training_args),
**vars(hypernet_args),
**vars(aggregator_args),
**vars(ctx_encoder_args),
}
run_name = os.path.basename(output_dir)
training_args.run_name = run_name
training_args.output_dir = output_dir
training_args.logging_dir = output_dir
training_args.save_safetensors = False
logger.debug(f"args: {args}")
save_yaml(args, f"{output_dir}/args.yaml")
############ Model setup
@ -230,19 +237,14 @@ def main(output_dir):
if ctx_args.exp_setup == ExperimentSetup.HYPER_LORA:
logger.info("Using HyperLoRA")
hypernet = HyperLoRA(
get_hypernet_config(model, hypernet_args, aggregator_args),
model,
).to(model.device)
# HACK: hardcode the embedding layer for now
# TODO: add ctx_encoder config
# ctx_encoder = torch.nn.Embedding.from_pretrained(
# model.get_input_embeddings().weight.clone(),
# freeze=True,
# )
# 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)
hypernet_config = get_hypernet_config(model, hypernet_args, aggregator_args)
# hypernet = HyperLoRA(
# get_hypernet_config(model, hypernet_args, aggregator_args),
# model,
# ).to(model.device)
# ctx_encoder = EarlyExit(get_base_model(model), ctx_encoder_args.layer_idx)
model = ModulatedPretrainedModel(model, hypernet_config, ctx_encoder_args)
else:
# activate LoRA
logger.info("Using LoRA")
@ -281,108 +283,6 @@ def main(output_dir):
[_get_tokenized_dataset(ds_name, split) for ds_name in ds_names]
)
# 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))
# # 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"]
val_train_indices = np.random.permutation(len(train_ds))[:500]
val_ds = {
@ -512,7 +412,7 @@ if __name__ == "__main__":
output_dir = f"train_outputs/{run_name}"
setup_logging(output_dir, debug=os.environ.get("DEBUG", False))
logger.debug(f"CMD: {' '.join(os.sys.argv)}")
save_yaml(extract_cli_args(os.sys.argv), f"{output_dir}/config.yaml")
save_yaml(extract_cli_args(os.sys.argv), f"{output_dir}/cli_args.yaml")
# disable_caching()
main(output_dir)

View file

@ -9,13 +9,13 @@ from typing import Any, Iterable, Optional, Tuple, Union
import torch
import torch.nn.functional as F
from configs import AggregatorArguments, HypernetArguments
from configs import AggregatorArguments, HypernetArguments, CtxEncoderArguments
from einops import rearrange, repeat, unpack
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 model_loading import get_lora_config, get_model, get_model_and_tokenizer
from peft import get_peft_config, load_peft_weights, LoraConfig, PeftConfig, PeftModel
from peft.tuners._buffer_dict import BufferDict
from peft.tuners.tuners_utils import BaseTunerLayer, check_target_module_exists
@ -26,7 +26,12 @@ from transformers.models.perceiver.modeling_perceiver import (
PerceiverBasicDecoder,
)
from transformers.modeling_outputs import ModelOutput
from utils import get_lora_module_names, get_num_layers, get_peft_in_out_features
from utils import (
get_lora_module_names,
get_num_layers,
get_peft_in_out_features,
get_base_model,
)
logger = logging.getLogger()
@ -39,7 +44,7 @@ class AGGREGATOR_TYPE(str, Enum):
@dataclass
class AggregatorConfig:
aggregator_type: AGGREGATOR_TYPE
# pooler
pooling_type: POOL_FN
feature_size: int
@ -361,41 +366,32 @@ def get_init_peft_weights(model: PeftModel, peft_config: PeftConfig = None):
class HyperLoRA(nn.Module):
def __init__(
self,
# latent_size: int,
# lora_config: LoraConfig,
# layer_indices: Integer[Tensor, "num_layers"],
# in_out_features: dict,
# aggregator_type: AGGREGATOR_TYPE,
# aggregator_kwargs: dict,
config: HypernetConfig,
base_model: Optional[PreTrainedModel] = None,
):
def __init__(self, config: HypernetConfig):
super().__init__()
# aggregator output [bs, n_layers, n_modules, feature_dim]
# by mixing the pooled features with layer embs and module embs (for pooling)
# or via a perceiver w/ bottleneck size = n_modules * n_layers
agg_config = config.aggregator_config
self.aggregator = AG[agg_config.aggregator_type](**vars(agg_config))
self.config = config
self._init_model()
self.lora_config = config.lora_config
def _init_model(self):
self.agg_config = self.config.aggregator_config
self.aggregator = AG[self.agg_config.aggregator_type](**vars(self.agg_config))
self.lora_config = self.config.lora_config
self.target_modules = self.lora_config.target_modules
self.layer_indices = config.layer_indices
self.layer_indices = self.config.layer_indices
self.in_d, self.out_d = config.feature_sizes
self.in_d, self.out_d = self.config.feature_sizes
self.layers = MLPResidualBlock(
input_size=config.latent_size,
hidden_size=config.latent_size * 4,
output_size=config.latent_size,
input_size=self.config.latent_size,
hidden_size=self.config.latent_size * 4,
output_size=self.config.latent_size,
)
# TODO: check initialization of the head
# default values are prob. way too big
#
# 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
@ -406,30 +402,10 @@ class HyperLoRA(nn.Module):
# bias_shape=None, # no bias
bias_shape="n_modules r out_d",
n_modules=len(self.target_modules),
d=config.latent_size,
r=config.lora_config.r,
d=self.config.latent_size,
r=self.config.lora_config.r,
out_d=max(self.in_d[m] + self.out_d[m] for m in self.target_modules),
)
if base_model is not None:
self._init_head(base_model)
@torch.no_grad()
def _init_head(self, base_model: PreTrainedModel):
peft_weights = get_init_peft_weights(base_model, self.lora_config)
logger.debug(f"peft_weights: {peft_weights}")
self.head.weight.data[:] = 0
self.head.bias.data[:] = 0
for i, m in enumerate(self.target_modules):
A = peft_weights[m]["lora_A"].weight.clone() # [r, in_d]
B = peft_weights[m]["lora_B"].weight.clone() # [out_d, r]
biases = [A, B.T]
# bias_hyper_init(self.head.bias[i], biases)
# bias-hyperinit
# init weights to zeros and bias to the base weights
bias_cat = torch.cat(biases, dim=1)
self.head.bias.data[..., i, :, : bias_cat.shape[1]] = bias_cat
self.head.bias.requires_grad = False
def _to_lora_dict(
self, flat_loras: Float[Tensor, "bs n_layers n_modules r max_io_dim"]
@ -489,16 +465,41 @@ class HyperLoRA(nn.Module):
class ModulatedPretrainedModel(nn.Module):
def __init__(
self,
base_model: PreTrainedModel, # TODO: maybe just use the name/config?
hypernet: HyperLoRA,
ctx_encoder: nn.Module, # TODO: maybe just use the name/config?
base_model: PreTrainedModel,
hypernet_config: HypernetConfig,
ctx_encoder_args: CtxEncoderArguments,
):
super().__init__()
self.device = base_model.device
# register base_model as a submodule
self.hypernet_config = hypernet_config
self.ctx_encoder_args = ctx_encoder_args
self.register_module("base_model", base_model)
self.register_module("hypernet", hypernet)
self.register_module("ctx_encoder", ctx_encoder)
self._init_model()
self._bias_hyper_init()
# self.register_module("hypernet", hypernet)
# self.register_module("ctx_encoder", ctx_encoder)
@classmethod
def from_state_dict(cls, state_dict: dict, train: bool = True):
lora_config = state_dict["hypernet_config"].lora_config
model_name_or_path = lora_config.base_model_name_or_path
base_model = get_model(
model_name_or_path,
train=train,
requires_grad=False,
peft_config=lora_config,
)
hypernet_config = state_dict.pop("hypernet_config")
ctx_encoder_args = state_dict.pop("ctx_encoder_args")
return cls(base_model, hypernet_config, ctx_encoder_args)
def _init_model(self):
self.hypernet = HyperLoRA(self.hypernet_config).to(self.device)
# TODO: allow ctx_encoder to be other models
self.ctx_encoder = EarlyExit(
get_base_model(self.base_model), self.ctx_encoder_args.layer_idx
)
# Delegate to base_model
@property
@ -508,22 +509,44 @@ class ModulatedPretrainedModel(nn.Module):
def get_input_embeddings(self):
return self.base_model.get_input_embeddings()
@torch.no_grad()
def _bias_hyper_init(self):
peft_weights = get_init_peft_weights(self.base_model, self.hypernet.lora_config)
logger.debug(f"peft_weights: {peft_weights}")
self.hypernet.head.weight.data[:] = 0
self.hypernet.head.bias.data[:] = 0
for i, m in enumerate(self.hypernet.target_modules):
A = peft_weights[m]["lora_A"].weight.clone() # [r, in_d]
B = peft_weights[m]["lora_B"].weight.clone() # [out_d, r]
biases = [A, B.T]
# bias-hyperinit
# init weights to zeros and bias to the base weights
bias_cat = torch.cat(biases, dim=1)
self.hypernet.head.bias.data[..., i, :, : bias_cat.shape[1]] = bias_cat
# self.hypernet.head.bias.requires_grad = False
def state_dict(self, *args, **kwargs):
state_dict = super().state_dict(*args, **kwargs)
# remove non-trainable and base_model's params
for name, param in self.named_parameters():
if not param.requires_grad:
state_dict.pop(name, None)
for name in list(state_dict.keys()):
if name.startswith("base_model") or name.startswith("ctx_encoder"):
state_dict.pop(name, None)
state_dict = self.hypernet.state_dict(*args, **kwargs)
state_dict["hypernet_config"] = self.hypernet_config
state_dict["ctx_encoder_args"] = self.ctx_encoder_args
return state_dict
def load_state_dict(self, state_dict: dict, *args, **kwargs):
# NOTE: might have to set `strict=False` as we don't save all the params
return super().load_state_dict(state_dict, *args, **kwargs)
self.hypernet_config = state_dict.pop("hypernet_config")
self.ctx_encoder_args = state_dict.pop("ctx_encoder_args")
if (
self.hypernet_config.lora_config.base_model_name_or_path
!= self.base_model.name_or_path
):
raise ValueError(
f"Base model name or path mismatch. "
f"The base model given is: {self.base_model.name_or_path}, "
f"but the hypernet config is for: {self.hypernet_config.lora_config.base_model_name_or_path}"
)
self._init_model()
return self.hypernet.load_state_dict(state_dict, *args, **kwargs)
# @torch.no_grad()
# def get_ctx_features(
@ -658,7 +681,6 @@ def apply_generated_loras(
if __name__ == "__main__":
from utils import get_base_model
# set torch randomness seed
torch.manual_seed(42)
@ -669,38 +691,42 @@ if __name__ == "__main__":
requires_grad=False,
peft_config=get_lora_config(model_name),
)
# TODO: base model shoul be init with target lora config
print(base_model)
device = base_model.device
# lora_config = base_model.peft_config["default"]
# in_d, out_d = get_peft_in_out_features(base_model, peft_config=lora_config)
hypernet_args = HypernetArguments(latent_size=512)
aggregator_args = AggregatorArguments(aggregator_type=AGGREGATOR_TYPE.PERCEIVER)
hypernet_config = get_hypernet_config(base_model, hypernet_args, aggregator_args)
ctx_encoder_args = CtxEncoderArguments(layer_idx=4)
# ctx_encoder = EarlyExit(get_base_model(base_model), 4)
# hypernet = HyperLoRA(
# lora_config,
# torch.arange(get_num_layers(base_model), device=base_model.device),
# in_out_features={"in": in_d, "out": out_d},
# aggregator_type=AGGREGATOR_TYPE.POOLER,
# aggregator_config=get_aggregator_config(base_model, POOL_FN.MEAN),
# ).to(base_model.device)
# get_hypernet_config(base_model, hypernet_args, aggregator_args),
# base_model,
# ).to(device)
ctx_encoder = EarlyExit(get_base_model(base_model), 4)
hypernet = HyperLoRA(get_hypernet_config(base_model), base_model).to(device)
model = ModulatedPretrainedModel(base_model, hypernet, ctx_encoder).to(device)
model = ModulatedPretrainedModel(base_model, hypernet_config, ctx_encoder_args).to(
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(device)
ctx_ids = ctx_inputs["input_ids"]
ctx_attn_mask = ctx_inputs["attention_mask"]
ctx_features = model.ctx_encoder(input_ids=ctx_ids, attention_mask=ctx_attn_mask)
ctx_features = model.ctx_encoder(input_ids=ctx_ids, attention_mask=ctx_attn_mask).to(
torch.float32
)
print(ctx_ids.shape)
model.eval()
agg_features = hypernet.aggregator(ctx_features, ctx_attn_mask)
agg_features = model.hypernet.aggregator(ctx_features, ctx_attn_mask)
print(agg_features)
print(agg_features.shape)
hnetout = hypernet(ctx_features, ctx_attn_mask)
hnetout = model.hypernet(ctx_features, ctx_attn_mask)
print(hnetout)
print(hnetout.shape)