mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
data exp + packing
This commit is contained in:
parent
cfd29ed761
commit
b61086034c
24 changed files with 458 additions and 103598 deletions
65
configs/data_exp/qa.yaml
Normal file
65
configs/data_exp/qa.yaml
Normal file
|
|
@ -0,0 +1,65 @@
|
|||
output_dir: "" # just a placeholder
|
||||
bf16: true
|
||||
model_name_or_path: google/gemma-2-2b-it
|
||||
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: 8
|
||||
per_device_eval_batch_size: 8
|
||||
max_val_samples_per_ds: 1000
|
||||
# optim: schedule_free_adamw
|
||||
|
||||
learning_rate: 0.00004
|
||||
# lr_scheduler_type: "constant_with_warmup"
|
||||
neftune_noise_alpha: 5
|
||||
weight_decay: 0.01
|
||||
#
|
||||
warmup_steps: 100
|
||||
|
||||
dataloader_prefetch_factor: 8
|
||||
dataloader_num_workers: 8
|
||||
# LoRA
|
||||
lora_r: 8
|
||||
lora_dropout: 0.0
|
||||
target_modules:
|
||||
- down_proj
|
||||
# data
|
||||
train_ds_names:
|
||||
- fw_qa_3_small # ~ 20M
|
||||
- ctx_qa # 300k
|
||||
- pwc # 240k
|
||||
- hotpot_qa # 90k
|
||||
- squad # 90k
|
||||
- drop # 77k
|
||||
- narrativeqa # 40k
|
||||
- quoref # 11k
|
||||
- ropes # 11k
|
||||
- synthetic_convqa # 40k
|
||||
|
||||
val_ds_names:
|
||||
- fw_qa_3_pretrain
|
||||
- self_gen/google/gemma-2-2b-it/fw_qa_3_small
|
||||
- self_gen/google/gemma-2-2b-it/ctx_qa
|
||||
- self_gen/google/gemma-2-2b-it/pwc
|
||||
- self_gen/google/gemma-2-2b-it/hotpot_qa
|
||||
- self_gen/google/gemma-2-2b-it/squad
|
||||
- fw_qa_3
|
||||
- ctx_qa
|
||||
- pwc
|
||||
- hotpot_qa
|
||||
- squad
|
||||
|
||||
load_best_model_at_end: false
|
||||
metric_for_best_model: eval_fw_qa_3_pretrain_loss
|
||||
66
configs/data_exp/self_gen_3_and_pretrain_small.yaml
Normal file
66
configs/data_exp/self_gen_3_and_pretrain_small.yaml
Normal file
|
|
@ -0,0 +1,66 @@
|
|||
output_dir: "" # just a placeholder
|
||||
bf16: true
|
||||
model_name_or_path: google/gemma-2-2b-it
|
||||
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: 8
|
||||
per_device_eval_batch_size: 8
|
||||
max_val_samples_per_ds: 1000
|
||||
# optim: schedule_free_adamw
|
||||
|
||||
learning_rate: 0.00004
|
||||
# lr_scheduler_type: "constant_with_warmup"
|
||||
neftune_noise_alpha: 5
|
||||
weight_decay: 0.01
|
||||
#
|
||||
warmup_steps: 100
|
||||
|
||||
dataloader_prefetch_factor: 8
|
||||
dataloader_num_workers: 8
|
||||
# LoRA
|
||||
lora_r: 8
|
||||
lora_dropout: 0.0
|
||||
target_modules:
|
||||
- down_proj
|
||||
# data
|
||||
train_ds_names:
|
||||
- fw_qa_3_small_pretrain # ~6M
|
||||
- self_gen/google/gemma-2-2b-it/fw_qa_3_small # ~20M
|
||||
- self_gen/google/gemma-2-2b-it/ctx_qa # 300k
|
||||
- self_gen/google/gemma-2-2b-it/pwc # 240k
|
||||
- self_gen/google/gemma-2-2b-it/hotpot_qa # 90k
|
||||
- self_gen/google/gemma-2-2b-it/squad # 90k
|
||||
- self_gen/google/gemma-2-2b-it/drop # 77k
|
||||
- self_gen/google/gemma-2-2b-it/narrativeqa # 40k
|
||||
- self_gen/google/gemma-2-2b-it/quoref # 11k
|
||||
- self_gen/google/gemma-2-2b-it/ropes # 11k
|
||||
- self_gen/google/gemma-2-2b-it/synthetic_convqa # 40k
|
||||
|
||||
val_ds_names:
|
||||
- fw_qa_3_pretrain
|
||||
- self_gen/google/gemma-2-2b-it/fw_qa_3_small
|
||||
- self_gen/google/gemma-2-2b-it/ctx_qa
|
||||
- self_gen/google/gemma-2-2b-it/pwc
|
||||
- self_gen/google/gemma-2-2b-it/hotpot_qa
|
||||
- self_gen/google/gemma-2-2b-it/squad
|
||||
- fw_qa_3
|
||||
- ctx_qa
|
||||
- pwc
|
||||
- hotpot_qa
|
||||
- squad
|
||||
|
||||
load_best_model_at_end: false
|
||||
metric_for_best_model: eval_fw_qa_3_pretrain_loss
|
||||
67
configs/data_exp/self_gen_3_and_pretrain_small_aug.yaml
Normal file
67
configs/data_exp/self_gen_3_and_pretrain_small_aug.yaml
Normal file
|
|
@ -0,0 +1,67 @@
|
|||
output_dir: "" # just a placeholder
|
||||
bf16: true
|
||||
model_name_or_path: google/gemma-2-2b-it
|
||||
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: 8
|
||||
per_device_eval_batch_size: 8
|
||||
max_val_samples_per_ds: 1000
|
||||
# optim: schedule_free_adamw
|
||||
|
||||
learning_rate: 0.00004
|
||||
# lr_scheduler_type: "constant_with_warmup"
|
||||
neftune_noise_alpha: 5
|
||||
weight_decay: 0.01
|
||||
#
|
||||
warmup_steps: 100
|
||||
|
||||
dataloader_prefetch_factor: 8
|
||||
dataloader_num_workers: 8
|
||||
# LoRA
|
||||
lora_r: 8
|
||||
lora_dropout: 0.0
|
||||
target_modules:
|
||||
- down_proj
|
||||
# data
|
||||
train_ds_names:
|
||||
- fw_qa_3_small_pretrain # ~6M
|
||||
- fw_qa_3_small_aug_pretrain # ~12M
|
||||
- self_gen/google/gemma-2-2b-it/fw_qa_3_small # ~20M
|
||||
- self_gen/google/gemma-2-2b-it/ctx_qa # 300k
|
||||
- self_gen/google/gemma-2-2b-it/pwc # 240k
|
||||
- self_gen/google/gemma-2-2b-it/hotpot_qa # 90k
|
||||
- self_gen/google/gemma-2-2b-it/squad # 90k
|
||||
- self_gen/google/gemma-2-2b-it/drop # 77k
|
||||
- self_gen/google/gemma-2-2b-it/narrativeqa # 40k
|
||||
- self_gen/google/gemma-2-2b-it/quoref # 11k
|
||||
- self_gen/google/gemma-2-2b-it/ropes # 11k
|
||||
- self_gen/google/gemma-2-2b-it/synthetic_convqa # 40k
|
||||
|
||||
val_ds_names:
|
||||
- fw_qa_3_pretrain
|
||||
- self_gen/google/gemma-2-2b-it/fw_qa_3_small
|
||||
- self_gen/google/gemma-2-2b-it/ctx_qa
|
||||
- self_gen/google/gemma-2-2b-it/pwc
|
||||
- self_gen/google/gemma-2-2b-it/hotpot_qa
|
||||
- self_gen/google/gemma-2-2b-it/squad
|
||||
- fw_qa_3
|
||||
- ctx_qa
|
||||
- pwc
|
||||
- hotpot_qa
|
||||
- squad
|
||||
|
||||
load_best_model_at_end: false
|
||||
metric_for_best_model: eval_fw_qa_3_pretrain_loss
|
||||
65
configs/data_exp/self_gen_3_small.yaml
Normal file
65
configs/data_exp/self_gen_3_small.yaml
Normal file
|
|
@ -0,0 +1,65 @@
|
|||
output_dir: "" # just a placeholder
|
||||
bf16: true
|
||||
model_name_or_path: google/gemma-2-2b-it
|
||||
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: 8
|
||||
per_device_eval_batch_size: 8
|
||||
max_val_samples_per_ds: 1000
|
||||
# optim: schedule_free_adamw
|
||||
|
||||
learning_rate: 0.00004
|
||||
# lr_scheduler_type: "constant_with_warmup"
|
||||
neftune_noise_alpha: 5
|
||||
weight_decay: 0.01
|
||||
#
|
||||
warmup_steps: 100
|
||||
|
||||
dataloader_prefetch_factor: 8
|
||||
dataloader_num_workers: 8
|
||||
# LoRA
|
||||
lora_r: 8
|
||||
lora_dropout: 0.0
|
||||
target_modules:
|
||||
- down_proj
|
||||
# data
|
||||
train_ds_names:
|
||||
- self_gen/google/gemma-2-2b-it/fw_qa_3_small # ~20M
|
||||
- self_gen/google/gemma-2-2b-it/ctx_qa # 300k
|
||||
- self_gen/google/gemma-2-2b-it/pwc # 240k
|
||||
- self_gen/google/gemma-2-2b-it/hotpot_qa # 90k
|
||||
- self_gen/google/gemma-2-2b-it/squad # 90k
|
||||
- self_gen/google/gemma-2-2b-it/drop # 77k
|
||||
- self_gen/google/gemma-2-2b-it/narrativeqa # 40k
|
||||
- self_gen/google/gemma-2-2b-it/quoref # 11k
|
||||
- self_gen/google/gemma-2-2b-it/ropes # 11k
|
||||
- self_gen/google/gemma-2-2b-it/synthetic_convqa # 40k
|
||||
|
||||
val_ds_names:
|
||||
- fw_qa_3_pretrain
|
||||
- self_gen/google/gemma-2-2b-it/fw_qa_3_small
|
||||
- self_gen/google/gemma-2-2b-it/ctx_qa
|
||||
- self_gen/google/gemma-2-2b-it/pwc
|
||||
- self_gen/google/gemma-2-2b-it/hotpot_qa
|
||||
- self_gen/google/gemma-2-2b-it/squad
|
||||
- fw_qa_3
|
||||
- ctx_qa
|
||||
- pwc
|
||||
- hotpot_qa
|
||||
- squad
|
||||
|
||||
load_best_model_at_end: false
|
||||
metric_for_best_model: eval_fw_qa_3_pretrain_loss
|
||||
|
|
@ -49,7 +49,7 @@ train_ds_names:
|
|||
- self_gen/google/gemma-2-2b-it/synthetic_convqa # 40k
|
||||
|
||||
val_ds_names:
|
||||
- self_gen/google/gemma-2-2b-it/fw_qa_3
|
||||
- self_gen/google/gemma-2-2b-it/fw_qa_3_small
|
||||
- self_gen/google/gemma-2-2b-it/fw_qa_xl
|
||||
- self_gen/google/gemma-2-2b-it/ctx_qa
|
||||
- self_gen/google/gemma-2-2b-it/pwc
|
||||
|
|
@ -64,3 +64,4 @@ val_ds_names:
|
|||
|
||||
load_best_model_at_end: true
|
||||
metric_for_best_model: eval_pwc_loss
|
||||
# metric_for_best_model: eval_fw_qa_3_loss
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
34
intx_sft.py
34
intx_sft.py
|
|
@ -229,6 +229,8 @@ def main():
|
|||
if ctx_name is None:
|
||||
ctx_name = model.base_model.config.name_or_path
|
||||
ctx_tokenizer = get_tokenizer(ctx_name, train=True)
|
||||
training_args.gen_lora_l1_reg_coef = ctx_args.gen_lora_l1_reg_coef
|
||||
|
||||
if len([p for p in model.ctx_encoder.parameters() if p.requires_grad]):
|
||||
raise ValueError("ctx_encoder contains trainable parameters")
|
||||
if len([p for p in model.base_model.parameters() if p.requires_grad]):
|
||||
|
|
@ -358,19 +360,25 @@ def main():
|
|||
seed=training_args.seed,
|
||||
stopping_strategy="all_exhausted",
|
||||
)
|
||||
logger.info(f"Train dataset length: {len(train_ds)}")
|
||||
|
||||
if ctx_args.use_sequence_packing:
|
||||
logging.info("Packing dataset")
|
||||
train_ds = pack(
|
||||
train_ds,
|
||||
ctx_args.max_packed_inp_len,
|
||||
ctx_args.max_packed_ctx_len,
|
||||
max_packed_size=-1,
|
||||
num_proc=8,
|
||||
)
|
||||
# TODO: add stats here
|
||||
logging.info("Setting per_device_train_batch_size to 1")
|
||||
training_args.per_device_train_batch_size = 1
|
||||
with training_args.main_process_first():
|
||||
logging.info("Packing dataset")
|
||||
old_ds_len = len(train_ds)
|
||||
train_ds = pack(
|
||||
train_ds,
|
||||
ctx_args.max_packed_inp_len,
|
||||
ctx_args.max_packed_ctx_len,
|
||||
max_packed_size=-1,
|
||||
num_proc=8,
|
||||
)
|
||||
logger.info(f"Train dataset length: {len(train_ds)}")
|
||||
logger.info(
|
||||
f"Avg. # of samples per packed sequence: {old_ds_len / len(train_ds)}"
|
||||
)
|
||||
logger.info("Setting per_device_train_batch_size to 1")
|
||||
training_args.per_device_train_batch_size = 1
|
||||
|
||||
logger.info(f"train_ds: {train_ds}")
|
||||
logger.info(f"val_ds: {val_ds}")
|
||||
|
|
@ -383,11 +391,13 @@ def main():
|
|||
)
|
||||
|
||||
# TODO: use SFTTrainer instead? https://huggingface.co/docs/trl/en/sft_trainer
|
||||
# TODO: use packing with SFTTrainer
|
||||
|
||||
# HACK [local patch]: deepspeed model loading problem (for resume training)
|
||||
# see https://github.com/microsoft/DeepSpeed/pull/6626/files
|
||||
# /home/rujikorn_sakana_ai/.conda/envs/ctx-to-lora/lib/python3.10/site-packages/deepspeed/runtime/engine.py
|
||||
# HACK: [local path]: make liger kernel works with gemma3
|
||||
# see .venv/lib/python3.10/site-packages/liger_kernel/transformers/monkey_patch.py
|
||||
# L810:814 + L781
|
||||
if training_args.use_liger_kernel and is_liger_kernel_available():
|
||||
from liger_kernel.transformers import _apply_liger_kernel_to_instance
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ readme = "README.md"
|
|||
requires-python = ">= 3.10"
|
||||
dependencies = [
|
||||
"transformers==4.51.3",
|
||||
"deepspeed==0.16.7",
|
||||
"deepspeed==0.17.1",
|
||||
"accelerate==1.6.0",
|
||||
"datasets==3.6.0",
|
||||
"setuptools",
|
||||
|
|
|
|||
|
|
@ -6,18 +6,20 @@
|
|||
#SBATCH --output=outputs/%x-%j.out
|
||||
#SBATCH --error=outputs/%x-%j.out
|
||||
|
||||
uv run accelerate launch --num_processes=4 --gradient_accumulation_steps=8 --gradient_clipping=1.0 \
|
||||
uv run accelerate launch --num_processes=4 --gradient_accumulation_steps=16 --gradient_clipping=1.0 \
|
||||
--gpu_ids all --main_process_port 29561 intx_sft.py configs/qa.yaml \
|
||||
--model_name_or_path=google/gemma-2-2b-it --num_train_epochs=4 --per_device_train_batch_size=32 \
|
||||
--gradient_accumulation_steps=8 --per_device_eval_batch_size=32 --exp_setup=hyper_lora --aggregator_type=perceiver \
|
||||
--model_name_or_path=google/gemma-2-2b-it --num_train_epochs=2.2 --per_device_train_batch_size=-1 \
|
||||
--gradient_accumulation_steps=16 --per_device_eval_batch_size=4 --exp_setup=hyper_lora --aggregator_type=perceiver \
|
||||
--target_modules=down_proj \
|
||||
--num_self_attends_per_block=4 --num_latent_factor=1 \
|
||||
--num_self_attends_per_block=8 --num_latent_factor=2 \
|
||||
--lora_r=8 \
|
||||
--eval_steps=5000 --save_steps=5000 --learning_rate=4e-5 --lora_dropout=0.0 \
|
||||
--eval_steps=1000000000000000 --do_eval=no --do_predict=no --eval_on_start=False \
|
||||
--save_steps=5000 --learning_rate=4e-5 --lora_dropout=0.0 \
|
||||
--neftune_noise_alpha=5 --use_light_weight_lora=False \
|
||||
--add_negative_prompt=False \
|
||||
--add_repeat_prompt=False \
|
||||
--use_sequence_packing=True --per_rank_gen=True \
|
||||
--use_sequence_packing=True --max_packed_inp_len=8192 --max_packed_ctx_len=16384 \
|
||||
--per_rank_gen=True \
|
||||
--per_layer_processing=True \
|
||||
--gen_lora_l1_reg_coef=0.1 \
|
||||
|
||||
|
|
|
|||
25
scripts/gemma_data_exp/gemma_qa_8_gpu.sh
Normal file
25
scripts/gemma_data_exp/gemma_qa_8_gpu.sh
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
#!/bin/bash
|
||||
#SBATCH --job-name=ctxlora_medium
|
||||
#SBATCH --partition=a3
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --gpus=8
|
||||
#SBATCH --output=outputs/%x-%j.out
|
||||
#SBATCH --error=outputs/%x-%j.out
|
||||
|
||||
uv run accelerate launch --num_processes=8 --gradient_accumulation_steps=8 --gradient_clipping=1.0 \
|
||||
--gpu_ids all --main_process_port 29561 intx_sft.py configs/qa.yaml \
|
||||
--model_name_or_path=google/gemma-2-2b-it --num_train_epochs=2.2 --per_device_train_batch_size=-1 \
|
||||
--gradient_accumulation_steps=8 --per_device_eval_batch_size=4 --exp_setup=hyper_lora --aggregator_type=perceiver \
|
||||
--target_modules=down_proj \
|
||||
--num_self_attends_per_block=8 --num_latent_factor=2 \
|
||||
--lora_r=16 \
|
||||
--eval_steps=1000000000000000 --do_eval=no --do_predict=no --eval_on_start=False \
|
||||
--save_steps=5000 --learning_rate=4e-5 --lora_dropout=0.0 \
|
||||
--neftune_noise_alpha=5 --use_light_weight_lora=False \
|
||||
--add_negative_prompt=False \
|
||||
--add_repeat_prompt=False \
|
||||
--use_sequence_packing=True --max_packed_inp_len=8192 --max_packed_ctx_len=16384 \
|
||||
--per_rank_gen=True \
|
||||
--per_layer_processing=True \
|
||||
--gen_lora_l1_reg_coef=0.1 \
|
||||
|
||||
|
|
@ -8,16 +8,17 @@
|
|||
|
||||
uv run accelerate launch --num_processes=4 --gradient_accumulation_steps=16 --gradient_clipping=1.0 \
|
||||
--gpu_ids all --main_process_port 29562 intx_sft.py configs/self_gen_3_small.yaml \
|
||||
--model_name_or_path=google/gemma-2-2b-it --num_train_epochs=3 --per_device_train_batch_size=8 \
|
||||
--model_name_or_path=google/gemma-2-2b-it --num_train_epochs=2.2 --per_device_train_batch_size=-1 \
|
||||
--gradient_accumulation_steps=16 --per_device_eval_batch_size=4 --exp_setup=hyper_lora --aggregator_type=perceiver \
|
||||
--target_modules=down_proj \
|
||||
--num_self_attends_per_block=4 --num_latent_factor=1 \
|
||||
--num_self_attends_per_block=8 --num_latent_factor=2 \
|
||||
--lora_r=8 \
|
||||
--eval_steps=5000 --save_steps=5000 --learning_rate=4e-5 --lora_dropout=0.0 \
|
||||
--eval_steps=1000000000000000 --do_eval=no --do_predict=no --eval_on_start=False \
|
||||
--save_steps=5000 --learning_rate=4e-5 --lora_dropout=0.0 \
|
||||
--neftune_noise_alpha=5 --use_light_weight_lora=False \
|
||||
--add_negative_prompt=False \
|
||||
--add_repeat_prompt=False \
|
||||
--use_sequence_packing=True --per_rank_gen=True \
|
||||
--use_sequence_packing=True --max_packed_inp_len=8192 --max_packed_ctx_len=16384 \
|
||||
--per_rank_gen=True \
|
||||
--per_layer_processing=True \
|
||||
--gen_lora_l1_reg_coef=0.1 \
|
||||
|
||||
|
|
|
|||
|
|
@ -11,13 +11,13 @@ uv run accelerate launch --num_processes=4 --gradient_accumulation_steps=16 --gr
|
|||
--model_name_or_path=google/gemma-2-2b-it --num_train_epochs=3 --per_device_train_batch_size=8 \
|
||||
--gradient_accumulation_steps=16 --per_device_eval_batch_size=4 --exp_setup=hyper_lora --aggregator_type=perceiver \
|
||||
--target_modules=down_proj \
|
||||
--num_self_attends_per_block=8 --num_latent_factor=1 \
|
||||
--num_self_attends_per_block=8 --num_latent_factor=2 \
|
||||
--lora_r=8 \
|
||||
--eval_steps=5000 --save_steps=5000 --learning_rate=4e-5 --lora_dropout=0.0 \
|
||||
--neftune_noise_alpha=5 --use_light_weight_lora=False \
|
||||
--add_negative_prompt=False \
|
||||
--add_repeat_prompt=False \
|
||||
--use_sequence_packing=True --max_packed_inp_len=16384 --max_packed_ctx_len=32768 \
|
||||
--use_sequence_packing=True --max_packed_inp_len=8192 --max_packed_ctx_len=16384 \
|
||||
--per_rank_gen=True \
|
||||
--per_layer_processing=True \
|
||||
--gen_lora_l1_reg_coef=0.1 \
|
||||
|
|
|
|||
|
|
@ -3,21 +3,24 @@
|
|||
#SBATCH --partition=a3
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --gpus=4
|
||||
# NO ! #SBATCH --begin=now+6hour # delay scheduling
|
||||
#SBATCH --output=outputs/%x-%j.out
|
||||
#SBATCH --error=outputs/%x-%j.out
|
||||
|
||||
uv run accelerate launch --num_processes=4 --gradient_accumulation_steps=32 --gradient_clipping=1.0 \
|
||||
--gpu_ids all --main_process_port 29563 intx_sft.py configs/self_gen_3_and_pretrain_small.yaml \
|
||||
--model_name_or_path=google/gemma-2-2b-it --num_train_epochs=2 --per_device_train_batch_size=4 \
|
||||
--gradient_accumulation_steps=32 --per_device_eval_batch_size=4 --exp_setup=hyper_lora --aggregator_type=perceiver \
|
||||
uv run accelerate launch --num_processes=4 --gradient_accumulation_steps=16 --gradient_clipping=1.0 \
|
||||
--gpu_ids all --main_process_port 29563 intx_sft.py configs/self_gen_3_small.yaml \
|
||||
--model_name_or_path=google/gemma-2-2b-it --num_train_epochs=2 --per_device_train_batch_size=-1 \
|
||||
--gradient_accumulation_steps=16 --per_device_eval_batch_size=4 --exp_setup=hyper_lora --aggregator_type=perceiver \
|
||||
--target_modules=down_proj \
|
||||
--num_self_attends_per_block=4 --num_latent_factor=1 \
|
||||
--num_self_attends_per_block=8 --num_latent_factor=2 \
|
||||
--lora_r=8 \
|
||||
--eval_steps=5000 --save_steps=5000 --learning_rate=4e-5 --lora_dropout=0.0 \
|
||||
--eval_steps=1000000000000000 --do_eval=no --do_predict=no --eval_on_start=False \
|
||||
--save_steps=5000 --learning_rate=4e-5 --lora_dropout=0.0 \
|
||||
--neftune_noise_alpha=5 --use_light_weight_lora=False \
|
||||
--add_negative_prompt=False \
|
||||
--add_repeat_prompt=False \
|
||||
--use_sequence_packing=True --per_rank_gen=True \
|
||||
--add_repeat_prompt=True --repeat_prob=0.1 \
|
||||
--use_sequence_packing=True --max_packed_inp_len=8192 --max_packed_ctx_len=16384 \
|
||||
--per_rank_gen=True \
|
||||
--per_layer_processing=True \
|
||||
--gen_lora_l1_reg_coef=0.1 \
|
||||
|
||||
|
|
@ -0,0 +1,26 @@
|
|||
#!/bin/bash
|
||||
#SBATCH --job-name=ctxlora_medium
|
||||
#SBATCH --partition=a3
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --gpus=4
|
||||
# NO ! #SBATCH --begin=now+6hour # delay scheduling
|
||||
#SBATCH --output=outputs/%x-%j.out
|
||||
#SBATCH --error=outputs/%x-%j.out
|
||||
|
||||
uv run accelerate launch --num_processes=4 --gradient_accumulation_steps=16 --gradient_clipping=1.0 \
|
||||
--gpu_ids all --main_process_port 29564 intx_sft.py configs/data_exp/self_gen_3_and_pretrain_small.yaml \
|
||||
--model_name_or_path=google/gemma-2-2b-it --num_train_epochs=1.5 \
|
||||
--gradient_accumulation_steps=16 --per_device_eval_batch_size=4 --exp_setup=hyper_lora --aggregator_type=perceiver \
|
||||
--target_modules=down_proj \
|
||||
--num_self_attends_per_block=8 --num_latent_factor=2 \
|
||||
--lora_r=8 \
|
||||
--eval_steps=1000000000000000 --do_eval=no --do_predict=no --eval_on_start=False \
|
||||
--save_steps=5000 --learning_rate=4e-5 --lora_dropout=0.0 \
|
||||
--neftune_noise_alpha=5 --use_light_weight_lora=False \
|
||||
--add_negative_prompt=False \
|
||||
--add_repeat_prompt=True --repeat_prob=0.1 \
|
||||
--use_sequence_packing=True --max_packed_inp_len=8192 --max_packed_ctx_len=16384 \
|
||||
--per_rank_gen=True \
|
||||
--per_layer_processing=True \
|
||||
--gen_lora_l1_reg_coef=0.1 \
|
||||
|
||||
|
|
@ -0,0 +1,26 @@
|
|||
#!/bin/bash
|
||||
#SBATCH --job-name=ctxlora_medium
|
||||
#SBATCH --partition=a3
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --gpus=4
|
||||
# NO ! #SBATCH --begin=now+6hour # delay scheduling
|
||||
#SBATCH --output=outputs/%x-%j.out
|
||||
#SBATCH --error=outputs/%x-%j.out
|
||||
|
||||
uv run accelerate launch --num_processes=4 --gradient_accumulation_steps=16 --gradient_clipping=1.0 \
|
||||
--gpu_ids all --main_process_port 29565 intx_sft.py configs/data_exp/self_gen_3_and_pretrain_small_aug.yaml \
|
||||
--model_name_or_path=google/gemma-2-2b-it --num_train_epochs=1 \
|
||||
--gradient_accumulation_steps=16 --per_device_eval_batch_size=4 --exp_setup=hyper_lora --aggregator_type=perceiver \
|
||||
--target_modules=down_proj \
|
||||
--num_self_attends_per_block=8 --num_latent_factor=2 \
|
||||
--lora_r=8 \
|
||||
--eval_steps=1000000000000000 --do_eval=no --do_predict=no --eval_on_start=False \
|
||||
--save_steps=5000 --learning_rate=4e-5 --lora_dropout=0.0 \
|
||||
--neftune_noise_alpha=5 --use_light_weight_lora=False \
|
||||
--add_negative_prompt=False \
|
||||
--add_repeat_prompt=True --repeat_prob=0.1 \
|
||||
--use_sequence_packing=True --max_packed_inp_len=8192 --max_packed_ctx_len=16384 \
|
||||
--per_rank_gen=True \
|
||||
--per_layer_processing=True \
|
||||
--gen_lora_l1_reg_coef=0.1 \
|
||||
|
||||
|
|
@ -190,7 +190,7 @@ class TrainingArguments(TrainingArguments):
|
|||
metadata={"help": "Number of warmup steps."},
|
||||
)
|
||||
eval_on_start: bool = field(
|
||||
default=True,
|
||||
default=False,
|
||||
metadata={"help": "Whether to evaluate on the start of training."},
|
||||
)
|
||||
eval_strategy: str = field(
|
||||
|
|
@ -243,6 +243,14 @@ class TrainingArguments(TrainingArguments):
|
|||
batch_eval_metrics: bool = field(
|
||||
default=True,
|
||||
)
|
||||
logging_first_step: bool = field(
|
||||
default=True,
|
||||
metadata={"help": "Whether to log the first step."},
|
||||
)
|
||||
ddp_timeout: int = field(
|
||||
default=2**20,
|
||||
metadata={"help": "Timeout for distributed data parallel training."},
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
|
|
@ -373,13 +381,13 @@ class DataArguments:
|
|||
metadata={"help": "Test dataset names."},
|
||||
)
|
||||
max_val_samples_per_ds: int | None = field(
|
||||
default=5000,
|
||||
default=1000,
|
||||
metadata={"help": "Maximum number of validation samples per dataset."},
|
||||
)
|
||||
max_test_samples_per_ds: int | None = field(
|
||||
default=1000,
|
||||
metadata={"help": "Maximum number of test samples per dataset."},
|
||||
)
|
||||
# max_test_samples_per_ds: int | None = field(
|
||||
# default=1000,
|
||||
# metadata={"help": "Maximum number of test samples per dataset."},
|
||||
# )
|
||||
|
||||
|
||||
@dataclass
|
||||
|
|
@ -416,6 +424,9 @@ class HypernetArguments:
|
|||
default=False,
|
||||
metadata={"help": "Whether to use token mixing block."},
|
||||
)
|
||||
num_pre_head_layers: int = field(
|
||||
default=4, metadata={"help": "# of layers before hypernet head"}
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
|
|
@ -424,6 +435,7 @@ class CtxEncoderArguments:
|
|||
default=None,
|
||||
metadata={"help": "Context encoder model name or path."},
|
||||
)
|
||||
# TODO: allow using activations from multiple layers?
|
||||
layer_idx: int | None = field(
|
||||
default=None,
|
||||
metadata={
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
import itertools
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from transformers.data import (
|
||||
|
|
@ -8,6 +10,17 @@ from transformers.data import (
|
|||
flattener = DataCollatorWithFlattening()
|
||||
|
||||
|
||||
def concat_batch(inp_list):
|
||||
sample = inp_list[0]
|
||||
out = [
|
||||
{
|
||||
k: list(itertools.chain.from_iterable([x[k] for x in inp_list]))
|
||||
for k in sample
|
||||
}
|
||||
]
|
||||
return out
|
||||
|
||||
|
||||
# def train_packed_collator(inp_list):
|
||||
# # no padding
|
||||
# packed_inputs = flattener(inp_list, return_tensors="pt")
|
||||
|
|
@ -28,8 +41,12 @@ flattener = DataCollatorWithFlattening()
|
|||
def flatten_if_not_packed(inp_list):
|
||||
# no padding
|
||||
sample = inp_list[0]
|
||||
n = len(sample)
|
||||
if "position_ids" in sample:
|
||||
return default_data_collator(inp_list, return_tensors="pt")
|
||||
if n == 1:
|
||||
return default_data_collator(inp_list, return_tensors="pt")
|
||||
elif n > 1:
|
||||
return default_data_collator(concat_batch(inp_list), return_tensors="pt")
|
||||
|
||||
packed_inputs = flattener(inp_list, return_tensors="pt")
|
||||
if "ctx_ids" in sample:
|
||||
|
|
@ -154,6 +171,10 @@ def eval_collator(inp_list, tokenizer):
|
|||
out["chat_ids"] = chat_ids
|
||||
out["chat_attn_mask"] = chat_attn_mask
|
||||
out["chat_labels"] = chat_labels
|
||||
# if "input_ids_len" in inp_list[0]:
|
||||
# out["input_ids_len"] = [x.pop("input_ids_len") for x in inp_list]
|
||||
# if "ctx_ids_len" in inp_list[0]:
|
||||
# out["ctx_ids_len"] = [x.pop("ctx_ids_len") for x in inp_list]
|
||||
return out
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -178,6 +178,13 @@ DS_KWARGS = {
|
|||
split="train",
|
||||
)
|
||||
),
|
||||
"fw_qa_3_small_aug_pretrain": dict(
|
||||
train=dict(
|
||||
path="parquet",
|
||||
data_files=glob("data/raw_datasets/fw_qa_aug_pretrain/000_*[!val].parquet"),
|
||||
split="train",
|
||||
)
|
||||
),
|
||||
"fw_qa_3_medium_pretrain": dict(
|
||||
train=dict(
|
||||
path="parquet",
|
||||
|
|
|
|||
|
|
@ -50,10 +50,10 @@ def pack_data_points_by_length(
|
|||
valid_ends = valid_ends_inp & valid_ends_ctx
|
||||
|
||||
# this should never happen?
|
||||
# if not np.any(valid_ends):
|
||||
# # Single item exceeds max_packed_inp_len, skip it
|
||||
# i += 1
|
||||
# continue
|
||||
if not np.any(valid_ends):
|
||||
# Single item exceeds max_packed_inp_len, skip it
|
||||
i += 1
|
||||
continue
|
||||
|
||||
# Find the last valid index
|
||||
max_valid_idx = i + np.where(valid_ends)[0][-1]
|
||||
|
|
|
|||
|
|
@ -113,6 +113,10 @@ def train_model(
|
|||
logger.info("Training with modulated model. Using CustomTrainer.")
|
||||
trainer_kwargs["gen_lora_l1_reg_coef"] = training_args.gen_lora_l1_reg_coef
|
||||
del training_args.gen_lora_l1_reg_coef
|
||||
if training_args.auto_find_batch_size:
|
||||
# set the batch size to some high number
|
||||
# which will be lowered by the Trainer
|
||||
training_args.per_device_train_batch_size = 128
|
||||
|
||||
trainer = trainer_cls(**trainer_kwargs)
|
||||
|
||||
|
|
|
|||
22
uv.lock
generated
22
uv.lock
generated
|
|
@ -705,7 +705,7 @@ dependencies = [
|
|||
requires-dist = [
|
||||
{ name = "accelerate", specifier = "==1.6.0" },
|
||||
{ name = "datasets", specifier = "==3.6.0" },
|
||||
{ name = "deepspeed", specifier = "==0.16.7" },
|
||||
{ name = "deepspeed", specifier = "==0.17.1" },
|
||||
{ name = "einops" },
|
||||
{ name = "fasttext-wheel" },
|
||||
{ name = "flask" },
|
||||
|
|
@ -821,7 +821,7 @@ wheels = [
|
|||
|
||||
[[package]]
|
||||
name = "deepspeed"
|
||||
version = "0.16.7"
|
||||
version = "0.17.1"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "einops" },
|
||||
|
|
@ -836,7 +836,7 @@ dependencies = [
|
|||
{ name = "torch" },
|
||||
{ name = "tqdm" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/06/b3/a3903de5c5b707170c5c27e1a40f4ef613f14d241bd84d8b151a2a8786f6/deepspeed-0.16.7.tar.gz", hash = "sha256:2a56eee6ad7b82decf37bb6ad8ec5e16dad96b647d79854132c715d42a78bd85", size = 1511314, upload-time = "2025-04-18T15:37:44.929Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/38/10/a7f63e086c1e1c12e290c98363c748ef5ddd6313fde739d2aeccd5ed0cd4/deepspeed-0.17.1.tar.gz", hash = "sha256:6d6e21796982b9e024f489e1c211666cc6c0be6e344751368610b9d2da285d6e", size = 1547985, upload-time = "2025-06-09T22:53:11.543Z" }
|
||||
|
||||
[[package]]
|
||||
name = "defusedxml"
|
||||
|
|
@ -2980,7 +2980,7 @@ name = "nvidia-cudnn-cu12"
|
|||
version = "9.1.0.70"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "nvidia-cublas-cu12", marker = "sys_platform == 'linux'" },
|
||||
{ name = "nvidia-cublas-cu12", marker = "platform_machine != 'aarch64' and sys_platform == 'linux'" },
|
||||
]
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/9f/fd/713452cd72343f682b1c7b9321e23829f00b842ceaedcda96e742ea0b0b3/nvidia_cudnn_cu12-9.1.0.70-py3-none-manylinux2014_x86_64.whl", hash = "sha256:165764f44ef8c61fcdfdfdbe769d687e06374059fbb388b6c89ecb0e28793a6f", size = 664752741, upload-time = "2024-04-22T15:24:15.253Z" },
|
||||
|
|
@ -2991,7 +2991,7 @@ name = "nvidia-cufft-cu12"
|
|||
version = "11.2.1.3"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "nvidia-nvjitlink-cu12", marker = "sys_platform == 'linux'" },
|
||||
{ name = "nvidia-nvjitlink-cu12", marker = "platform_machine != 'aarch64' and sys_platform == 'linux'" },
|
||||
]
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/27/94/3266821f65b92b3138631e9c8e7fe1fb513804ac934485a8d05776e1dd43/nvidia_cufft_cu12-11.2.1.3-py3-none-manylinux2014_x86_64.whl", hash = "sha256:f083fc24912aa410be21fa16d157fed2055dab1cc4b6934a0e03cba69eb242b9", size = 211459117, upload-time = "2024-04-03T20:57:40.402Z" },
|
||||
|
|
@ -3010,9 +3010,9 @@ name = "nvidia-cusolver-cu12"
|
|||
version = "11.6.1.9"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "nvidia-cublas-cu12", marker = "sys_platform == 'linux'" },
|
||||
{ name = "nvidia-cusparse-cu12", marker = "sys_platform == 'linux'" },
|
||||
{ name = "nvidia-nvjitlink-cu12", marker = "sys_platform == 'linux'" },
|
||||
{ name = "nvidia-cublas-cu12", marker = "platform_machine != 'aarch64' and sys_platform == 'linux'" },
|
||||
{ name = "nvidia-cusparse-cu12", marker = "platform_machine != 'aarch64' and sys_platform == 'linux'" },
|
||||
{ name = "nvidia-nvjitlink-cu12", marker = "platform_machine != 'aarch64' and sys_platform == 'linux'" },
|
||||
]
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/3a/e1/5b9089a4b2a4790dfdea8b3a006052cfecff58139d5a4e34cb1a51df8d6f/nvidia_cusolver_cu12-11.6.1.9-py3-none-manylinux2014_x86_64.whl", hash = "sha256:19e33fa442bcfd085b3086c4ebf7e8debc07cfe01e11513cc6d332fd918ac260", size = 127936057, upload-time = "2024-04-03T20:58:28.735Z" },
|
||||
|
|
@ -3023,7 +3023,7 @@ name = "nvidia-cusparse-cu12"
|
|||
version = "12.3.1.170"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "nvidia-nvjitlink-cu12", marker = "sys_platform == 'linux'" },
|
||||
{ name = "nvidia-nvjitlink-cu12", marker = "platform_machine != 'aarch64' and sys_platform == 'linux'" },
|
||||
]
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/db/f7/97a9ea26ed4bbbfc2d470994b8b4f338ef663be97b8f677519ac195e113d/nvidia_cusparse_cu12-12.3.1.170-py3-none-manylinux2014_x86_64.whl", hash = "sha256:ea4f11a2904e2a8dc4b1833cc1b5181cde564edd0d5cd33e3c168eff2d1863f1", size = 207454763, upload-time = "2024-04-03T20:58:59.995Z" },
|
||||
|
|
@ -5587,8 +5587,8 @@ name = "xformers"
|
|||
version = "0.0.29.post2"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "numpy", marker = "sys_platform == 'linux'" },
|
||||
{ name = "torch", marker = "sys_platform == 'linux'" },
|
||||
{ name = "numpy", marker = "platform_machine != 'aarch64' and sys_platform == 'linux'" },
|
||||
{ name = "torch", marker = "platform_machine != 'aarch64' and sys_platform == 'linux'" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/27/ed/04ec7ef97a7e1c836add41ef5a2aef8cbdd45c0190ca42cc08f3c21e2b7b/xformers-0.0.29.post2.tar.gz", hash = "sha256:6ca3d1a6db6f2abff25c1154adee96987f77f4dfd5141771805afa5fc13e9395", size = 8468494, upload-time = "2025-02-01T02:33:48.209Z" }
|
||||
wheels = [
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue