This commit is contained in:
51616 2024-12-20 10:49:02 +00:00
parent 98d4b56629
commit bc7b588775
3 changed files with 11 additions and 2 deletions

View file

@ -1,4 +1,3 @@
output_dir: train_outputs/
bf16: true
target_modules:
- down_proj

View file

@ -43,7 +43,6 @@ class ArgumentParser(HfArgumentParser):
# overwrite the default/loaded value with the value provided to the command line
# adapted from https://github.com/huggingface/transformers/blob/d0b5002378daabf62769159add3e7d66d3f83c3b/src/transformers/hf_argparser.py#L327
print(arg_list)
for data_yaml, data_class in zip(arg_list, self.dataclass_types):
keys = {f.name for f in dataclasses.fields(data_yaml) if f.init}
inputs = {k: v for k, v in vars(data_yaml).items() if k in keys}

View file

@ -1,4 +1,7 @@
import logging
import random
import string
import time
import numpy as np
import torch
@ -64,6 +67,14 @@ def main():
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.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