data exp + packing

This commit is contained in:
51616 2025-06-17 11:34:16 +00:00
parent cfd29ed761
commit b61086034c
24 changed files with 458 additions and 103598 deletions

65
configs/data_exp/qa.yaml Normal file
View 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

View 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

View 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

View 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

View file

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

View file

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

View file

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

View file

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

View 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 \

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

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