mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
ctx_encoder_args + loadable modulated model + ctx_numbers_256
This commit is contained in:
parent
831e057eca
commit
c0cf32c8af
5 changed files with 690 additions and 201 deletions
555
configs/context_numbers_256.yaml
Normal file
555
configs/context_numbers_256.yaml
Normal 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
|
||||
|
||||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue