()
| 20 | |
| 21 | |
| 22 | def main(): |
| 23 | # torch.autograd.set_detect_anomaly(True) |
| 24 | args = get_train_args() |
| 25 | logger.info(args) |
| 26 | |
| 27 | tokenizer = AutoTokenizer.from_pretrained( |
| 28 | args.tokenizer, |
| 29 | use_fast=args.use_fast_tokenizer, |
| 30 | trust_remote_code=True, |
| 31 | add_bos_token=True, |
| 32 | add_eos_token=False |
| 33 | ) |
| 34 | if tokenizer.pad_token_id is None: |
| 35 | tokenizer.pad_token = tokenizer.eos_token |
| 36 | logger.info("Add pad token: {}".format(tokenizer.pad_token)) |
| 37 | # args.from_config = False |
| 38 | if args.from_config: |
| 39 | logger.info("All model params are randomly initialized for from-scratch training.") |
| 40 | model = AutoModelForCausalLM.from_config(AutoConfig.from_pretrained(args.model_name_or_path)) |
| 41 | else: |
| 42 | logger.info(f"Loading pretrained checkpoint {args.model_name_or_path}") |
| 43 | model = AutoModelForCausalLM.from_pretrained(args.model_name_or_path) |
| 44 | for name, param in model.named_parameters(): |
| 45 | if 'gate' in name: |
| 46 | if 'weight' in name: |
| 47 | nn.init.xavier_normal_(param) |
| 48 | model.train() |
| 49 | |
| 50 | # summary(model, depth=6) |
| 51 | # exit(0) |
| 52 | |
| 53 | trainable_params, all_param = model.num_parameters(only_trainable=True), model.num_parameters() |
| 54 | logger.info(f"% of trainable params: {trainable_params:d} / {all_param:d} = {trainable_params / all_param:.2%}") |
| 55 | logger.info(f"{tokenizer}\n{model}\n{model.config}") |
| 56 | |
| 57 | logger.info(f"Loading the `{args.split}` split directly from the cache {args.cache_dir}...") |
| 58 | dataset = load_from_disk(args.cache_dir) |
| 59 | logger.info(f"{dataset}") |
| 60 | logger.info(f"Shuffling the dataset with seed {args.seed}") |
| 61 | dataset = dataset.shuffle(seed=args.seed) |
| 62 | logger.info("Creating the data collator") |
| 63 | data_collator = DataCollatorForLanguageModeling(tokenizer=tokenizer, varlen=args.varlen) |
| 64 | logger.info(f"{data_collator}") |
| 65 | |
| 66 | if args.lr_scheduler_type == 'cosine_with_min_lr': |
| 67 | args.lr_scheduler_kwargs = {'min_lr_rate': 0.1} |
| 68 | if args.lr_scheduler_type == 'warmup_stable_decay': |
| 69 | args.lr_scheduler_kwargs = { |
| 70 | 'num_stable_steps': args.max_steps * 0.9 - args.warmup_steps, |
| 71 | 'num_decay_steps': args.max_steps * 0.1 |
| 72 | } |
| 73 | |
| 74 | args.logging_steps = 16 |
| 75 | trainer = Trainer( |
| 76 | model=model, |
| 77 | args=args, |
| 78 | tokenizer=tokenizer, |
| 79 | data_collator=data_collator, |
no test coverage detected