MCPcopy Create free account
hub / github.com/FireRedTeam/FireRedTTS2 / train

Function train

bin/finetune_example/posttrain.py:22–206  ·  view source on GitHub ↗

trial is only used when we are sweeping hyperparameters.

(args: argparse.Namespace, config: dict, trial: optuna.Trial = None)

Source from the content-addressed store, hash-verified

20
21
22def train(args: argparse.Namespace, config: dict, trial: optuna.Trial = None):
23 """
24 trial is only used when we are sweeping hyperparameters.
25 """
26
27 # accelerator
28 accelerator = Accelerator()
29 current_gpu = int(torch.cuda.current_device())
30 device = accelerator.device # "cuda"
31 n_gpus = torch.cuda.device_count()
32 print(
33 f"---Number of GPUs: {n_gpus}",
34 "---current_gpu:",
35 current_gpu,
36 "---device:",
37 device,
38 )
39
40 # prepare log
41 logs_folder = config["train"]["logs_folder"]
42 if accelerator.is_main_process:
43 writer = SummaryWriter(log_dir=logs_folder)
44
45 print("---Load LLM Model...")
46 model = load_model(config, args.checkpoint_path, device)
47
48 trainloader, valloader = create_dataloaders(
49 train_datasets=config["dataset"]["train_dataset_dir"],
50 validation_datasets=config["dataset"]["valid_dataset_dir"],
51 batch_size=config["train"]["batch_size"],
52 device=device,
53 infinite_train=False,
54 num_workers=8,
55 )
56
57 eff_batch_size = config["train"]["batch_size"] * config["train"]["accumulate_num"]
58
59 total_steps = (config["train"]["n_epochs"] * len(trainloader)) // config["train"][
60 "accumulate_num"
61 ]
62 print("---total_steps:", total_steps)
63
64 optimizer = torch.optim.AdamW(
65 model.parameters(),
66 lr=config["train"]["lr"],
67 weight_decay=config["train"]["weight_decay"],
68 )
69 scheduler = WarmupDecayLR(
70 optimizer,
71 config["train"]["warmup_steps"],
72 total_steps,
73 config["train"]["lr_decay"],
74 )
75
76 state = {
77 "model": model.state_dict(),
78 "optimizer": optimizer.state_dict(),
79 "scheduler": scheduler.state_dict(),

Callers 1

posttrain.pyFile · 0.85

Calls 5

load_modelFunction · 0.90
create_dataloadersFunction · 0.90
WarmupDecayLRClass · 0.90
get_grad_normFunction · 0.90
summarizeFunction · 0.90

Tested by

no test coverage detected