MCPcopy Create free account
hub / github.com/OpenSparseLLMs/MoM / main

Function main

training/run.py:22–100  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

20
21
22def 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,

Callers 1

run.pyFile · 0.70

Calls 4

get_train_argsFunction · 0.90
LogCallbackClass · 0.90
detect_nan_hookFunction · 0.85

Tested by

no test coverage detected