mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
runname
This commit is contained in:
parent
98d4b56629
commit
bc7b588775
3 changed files with 11 additions and 2 deletions
|
|
@ -1,4 +1,3 @@
|
|||
output_dir: train_outputs/
|
||||
bf16: true
|
||||
target_modules:
|
||||
- down_proj
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue