ScanDL2 / scandl_module /scripts /sp_run_train.py
Di0nigi's picture
First commit
95456ed verified
Raw
History Blame Contribute Delete
6.98 kB
import sys
import os
import argparse
import datetime
import time
from ScanDL2.CONSTANTS import (
COMPLETE_SCANDL_MODULE_TRAIN_PATH_BSC,
COMPLETE_SCANDL_MODULE_TRAIN_PATH_CELER,
COMPLETE_SCANDL_MODULE_TRAIN_PATH_EMTEC,
)
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="training args.")
parser.add_argument(
"--noise_schedule",
type=str,
default="sqrt",
choices=["linear", "cosine", "sqrt", "trunc_cos", "trunc_lin", "pw_lin"],
help="the distribution of noises",
)
parser.add_argument("--diff_steps", type=int, default=2000, help="diffusion steps")
parser.add_argument(
"--schedule_sampler",
type=str,
default="lossaware",
choices=["uniform", "lossaware", "fixstep"],
help="schedule sampler of timesteps",
)
parser.add_argument("--seq_len", type=int, default=128, help="max len of input sequence")
parser.add_argument(
"--hidden_t_dim", type=int, default=128, help="hidden size of time embedding"
)
parser.add_argument(
"--hidden_dim",
type=int,
default=768,
help="hidden size of word embedding and transformer hidden size",
)
parser.add_argument("--learning_steps", type=int, default=60000, help="total steps of learning")
parser.add_argument("--save_interval", type=int, default=2000, help="save step")
parser.add_argument(
"--resume_checkpoint",
type=str,
default="none",
help="path to resume checkpoint, like xxx/xxx.pt",
)
parser.add_argument("--lr", type=float, default=1e-04, help="learning rate")
parser.add_argument("--bsz", type=int, default=64, help="batch size")
parser.add_argument("--microbatch", type=int, default=64, help="microbatch size")
parser.add_argument("--seed", type=int, default=101, help="random seed")
parser.add_argument(
"--config_name", type=str, default="bert-base-cased", help="config of pre-trained models"
)
parser.add_argument(
"--vocab",
type=str,
default="bert",
help="use bert vocab or load external vocab dict if given as path",
)
parser.add_argument(
"--use_plm_init",
type=str,
default="no",
choices=["no", "bert"],
help="load init parameter from the pre-trained lm",
)
parser.add_argument("--log_interval", type=int, default=200, required=False)
parser.add_argument("--eval_interval", type=int, default=500, required=False)
parser.add_argument(
"--notes",
type=str,
default="-",
help="as training notes or specifical args",
required=False,
)
parser.add_argument("--app", type=str, default="", help="other input args")
# further arguments
parser.add_argument(
"--data_split_criterion",
type=str,
help="how to split the data into train, val, test:"
" scanpath (random), reader, sentence, combined",
required=False,
default="reader",
)
parser.add_argument(
"--num_transformer_layers",
type=int,
default=4,
required=False,
help="the number of encoder layers",
)
parser.add_argument(
"--num_transformer_heads",
type=int,
default=8,
required=False,
help="the number of attention heads",
)
parser.add_argument(
"--celer_only_L1",
required=False,
action="store_true",
help="if given, all celer speakers are used" "as opposed to only L1 speakers",
)
parser.add_argument(
"--corpus",
type=str,
help="the eye-tracking corpus to use for training.",
required=False,
default="celer",
choices=["celer", "zuco", "emtec", "bsc"],
)
parser.add_argument(
"--inference",
required=False,
default="cv",
choices=["cv", "zuco", "in-corpus"],
help="if zuco, inference is performed on zuco while trained on celer; if cv, inference is"
"done in k-fold Cross-Validation; if in-corpus, the training corpus is simply split into"
"train and test.",
)
parser.add_argument(
"--mask_padding",
action="store_false",
required=False,
help="if given, padding will not be masked in transformer attention. if not given, mask_padding"
"is stored as True; padding will be masked.",
)
parser.add_argument(
"--load_train_data",
type=str,
default="-",
help="if given, previously saved train data is loaded from the specified checkpoint path",
)
args = parser.parse_args()
# set working dir to the upper folder
abspath = os.path.abspath(sys.argv[0])
dname = os.path.dirname(abspath)
dname = os.path.dirname(dname)
os.chdir(dname)
if args.corpus == "emtec":
model_file = COMPLETE_SCANDL_MODULE_TRAIN_PATH_EMTEC
elif args.corpus == "bsc":
model_file = COMPLETE_SCANDL_MODULE_TRAIN_PATH_BSC
elif args.corpus == "celer":
model_file = COMPLETE_SCANDL_MODULE_TRAIN_PATH_CELER
else:
raise NotImplementedError(f"Corpus {args.corpus} not implemented.")
if int(os.environ["LOCAL_RANK"]) == 0:
if not os.path.exists(model_file):
os.makedirs(model_file)
COMMANDLINE = (
f"TOKENIZERS_PARALLELISM=FALSE "
f"python -m scripts.sp_train "
f"--checkpoint_path {model_file} "
f"--vocab {args.vocab} "
f"--use_plm_init {args.use_plm_init} "
f"--lr {args.lr} "
f"--batch_size {args.bsz} "
f"--microbatch {args.microbatch} "
f"--diffusion_steps {args.diff_steps} "
f"--noise_schedule {args.noise_schedule} "
f"--schedule_sampler {args.schedule_sampler} "
f"--seq_len {args.seq_len} "
f"--resume_checkpoint {args.resume_checkpoint} "
f"--hidden_t_dim {args.hidden_t_dim} "
f"--seed {args.seed} "
f"--hidden_dim {args.hidden_dim} "
f"--learning_steps {args.learning_steps} "
f"--save_interval {args.save_interval} "
f"--config_name {args.config_name} "
f"--notes {args.notes} "
f"--data_split_criterion {args.data_split_criterion} "
f"--num_transformer_layers {args.num_transformer_layers} "
f"--num_transformer_heads {args.num_transformer_heads} "
f"--corpus {args.corpus} "
f"--inference {args.inference} "
f"--load_train_data {args.load_train_data}"
)
if int(os.environ["LOCAL_RANK"]) == 0:
with open(os.path.join(model_file, "saved_bash.sh"), "w") as f:
print(COMMANDLINE, file=f)
print(COMMANDLINE)
os.system(COMMANDLINE)