refactor config + schedule free adam + liger kernel

This commit is contained in:
51616 2024-12-21 12:39:03 +00:00
parent c312d4bafb
commit 7980420114
6 changed files with 43 additions and 22 deletions

View file

@ -1,6 +1,25 @@
output_dir: "" # just a placeholder
bf16: true
base_model_name_or_path: meta-llama/Llama-3.2-1B-Instruct
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
optim: schedule_free_adamw
learning_rate: 0.00001
neftune_noise_alpha: 5
# LoRA
lora_r: 8
lora_dropout: 0.05
target_modules:
- down_proj
- up_proj
- gate_proj
- gate_proj

View file

@ -137,7 +137,7 @@ class ModelArguments:
@dataclass
class LoRAArguments:
r: Optional[int] = field(
lora_r: Optional[int] = field(
default=8,
metadata={"help": ("LoRA R value.")},
)

View file

@ -63,31 +63,15 @@ def main():
)
ctx_args, model_args, lora_args, training_args = parser.parse()
training_args.label_names = ["labels"]
training_args.eval_on_start = True
training_args.eval_strategy = "steps"
training_args.eval_steps = 500
training_args.save_strategy = "no"
# training_args.save_steps = 500
training_args.logging_strategy = "steps"
training_args.logging_steps = 100
uuid = "".join(
[random.choice(string.ascii_letters + string.digits) for _ in range(8)]
)
run_name = time.strftime("%Y%m%d-%H%M%S") + f"_{uuid}"
training_args.run_name = run_name
training_args.output_dir = f"train_outputs/{run_name}"
training_args.logging_dir = f"train_outputs/{run_name}"
# seq2seq args for generation evaluation
# training_args.predict_with_generate = True
# training_args.generation_max_length = 100
training_args.gradient_checkpointing_kwargs = {
"use_reentrant": False
} # manually add this argument in the code
model_name = model_args.model_name_or_path
model, tokenizer = get_model_and_tokenizer(

View file

@ -128,13 +128,13 @@ def get_model(
def get_lora_config(model_dir, **kwargs):
r = kwargs.get("r", 8)
r = kwargs.pop("lora_r", 8)
peft_conf_kwargs = dict(
r=r,
peft_type=PeftType.LORA,
base_model_name_or_path=model_dir,
task_type="CAUSAL_LM",
lora_dropout=0.05,
lora_dropout=kwargs.get("lora_dropout", 0.05),
lora_alpha=r ** (3 / 2) * 2,
)

14
install.sh Executable file
View file

@ -0,0 +1,14 @@
# install unsloth
# see https://docs.unsloth.ai/get-started/install-update/conda-install
conda create --name unsloth_env \
python=3.10 \
pytorch-cuda=<11.8/12.1> \
pytorch cudatoolkit xformers -c pytorch -c nvidia -c xformers \
-y
conda activate unsloth_env
pip install "unsloth[colab-new] @ git+https://github.com/unslothai/unsloth.git"
pip install --no-deps "trl<0.9.0" peft accelerate bitsandbytes
pip install -r requirements.txt

4
requirements.txt Normal file
View file

@ -0,0 +1,4 @@
einops
jaxtyping
schedulefree
liger-kernel