From c0cf32c8af475f0dfbbdc077bfd6450572e058f5 Mon Sep 17 00:00:00 2001 From: 51616 Date: Sun, 5 Jan 2025 10:35:43 +0000 Subject: [PATCH] ctx_encoder_args + loadable modulated model + ctx_numbers_256 --- configs/context_numbers_256.yaml | 555 +++++++++++++++++++++++++++++ data/raw_datasets/generate_data.py | 12 +- hyperlora/configs.py | 8 + hyperlora/intx_sft.py | 132 +------ hyperlora/modeling_utils.py | 184 ++++++---- 5 files changed, 690 insertions(+), 201 deletions(-) create mode 100644 configs/context_numbers_256.yaml diff --git a/configs/context_numbers_256.yaml b/configs/context_numbers_256.yaml new file mode 100644 index 0000000..6cd940c --- /dev/null +++ b/configs/context_numbers_256.yaml @@ -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 + diff --git a/data/raw_datasets/generate_data.py b/data/raw_datasets/generate_data.py index e6f48a3..58b4c49 100644 --- a/data/raw_datasets/generate_data.py +++ b/data/raw_datasets/generate_data.py @@ -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}" diff --git a/hyperlora/configs.py b/hyperlora/configs.py index 1a4731f..28a4d74 100644 --- a/hyperlora/configs.py +++ b/hyperlora/configs.py @@ -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: diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index b7f5c42..c2928d9 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -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) diff --git a/hyperlora/modeling_utils.py b/hyperlora/modeling_utils.py index ec226ad..9e8e10f 100644 --- a/hyperlora/modeling_utils.py +++ b/hyperlora/modeling_utils.py @@ -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)