import torch import json from torch.optim.lr_scheduler import LambdaLR import wandb from data import Vocab from data_gen import DataGenerator, Dataset, Augmenter from train_target import train_target_wants_text from model import ( ModelArgs, MusicalPositionEmbedTransformer, LabelSmoothing, SimpleLossCompute, RMSNorm, ) from main import TrainState, run_epoch, rate from embedder import Embedder import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP import dataclasses import logging import transformers.models.t5.modeling_t5 as t5 import dataset_classes import data_aug from pathlib import Path def run(config_path): with open(config_path, "r") as f: config = json.load(f) dist.init_process_group("nccl") rank = dist.get_rank() world_size = dist.get_world_size() device_id = rank % torch.cuda.device_count() torch.cuda.set_device(device_id) logging.info(f"start rank {rank+1} of {world_size}") vocab = Vocab.from_config(config) train_target = str(config["train_target"]) enable_text_prompt = train_target_wants_text(train_target) batch_size = int(config["batch_size"]) seq_len = int(config["seq_len"]) seq_len_min = int(config["seq_len_min"]) seq_len_max = int(config["seq_len_max"]) if enable_text_prompt: encoder = Embedder() else: encoder = None # construct dataset and generators from config train_split_idx = int(config["train_split_idx"]) eval_split_idx = int(config["eval_split_idx"]) example_continuous = bool(config["example_continuous"]) pack_batch = bool(config["pack_batch"]) dataset = Dataset.dynamic_from_config(config, "dataset_class") augmenter = Augmenter.dynamic_from_config(vocab, config, "augmenter_class") train_data_gen = DataGenerator( dataset, vocab, split=train_split_idx, ranksize=(rank, world_size), batch_size=batch_size, seq_len=seq_len, seq_len_min=seq_len_min, seq_len_max=seq_len_max, train_target=train_target, parallelism=8, to_device=f"cuda:{device_id}", text_tokenize=encoder.tokenize if enable_text_prompt else None, example_continuous=example_continuous, pack_batch=pack_batch, augmenter=augmenter, ) eval_data_gen = DataGenerator( dataset, vocab, split=eval_split_idx, ranksize=(rank, world_size), batch_size=batch_size, seq_len=seq_len, seq_len_min=seq_len_min, seq_len_max=seq_len_max, train_target=train_target, parallelism=8, to_device=f"cuda:{device_id}", text_tokenize=encoder.tokenize if enable_text_prompt else None, example_continuous=example_continuous, pack_batch=False, augmenter=augmenter, ) criterion = LabelSmoothing( padding_idx=0, smoothing=float(config["label_smoothing"]) ) model = MusicalPositionEmbedTransformer( vocab, dataclasses.replace( ModelArgs.from_config(vocab, config), cache=False, ), encoder, ) if rank == 0: print(model) print(f"model size: {model.param_count:.2e} parameters") # create model and move it to GPU with id rank model = model.to(device_id) model = DDP(model, device_ids=[device_id]) lr = 0.1 lr_factor = float(config["lr_factor"]) num_batches = torch.zeros(2, dtype=torch.long, device="cuda") if rank == 0: train_split_size = train_data_gen.num_examples() eval_split_size = eval_data_gen.num_examples() print( f"train split size {train_split_size:.2e}, eval split size {eval_split_size:.2e}" ) num_batches[0] = train_split_size // (batch_size * world_size) num_batches[1] = eval_split_size // (batch_size * world_size) assert num_batches[0] > 0, "train split is too small" assert num_batches[1] > 0, "eval split is too small" dist.broadcast(num_batches, 0) num_batches, num_eval_batches = num_batches.tolist() accum_iter = int(config["accum_iter"]) eval_iter = int(config["eval_iter"]) epochs = int(config["epochs"]) lr_decay_epochs = int(config["lr_decay_epochs"]) # choose whether to weight decay each module # lifted from minGPT decay = set() no_decay = set() whitelist_weight_modules = (torch.nn.Linear,) blacklist_weight_modules = ( torch.nn.LayerNorm, torch.nn.Embedding, RMSNorm, t5.T5LayerNorm, ) for mn, m in model.named_modules(): for pn, p in m.named_parameters(): fpn = "%s.%s" % (mn, pn) if mn else pn # full param name if pn.endswith("bias") or pn.endswith("alpha"): no_decay.add(fpn) elif pn.endswith("weight") and isinstance(m, whitelist_weight_modules): decay.add(fpn) elif pn.endswith("weight") and isinstance(m, blacklist_weight_modules): no_decay.add(fpn) # validate that we considered every parameter param_dict = {pn: p for pn, p in model.named_parameters()} inter_params = decay & no_decay union_params = decay | no_decay assert ( len(inter_params) == 0 ), "parameters %s made it into both decay/no_decay sets!" % (str(inter_params),) assert ( len(param_dict.keys() - union_params) == 0 ), "parameters %s were not separated into either decay/no_decay set!" % ( str(param_dict.keys() - union_params), ) # create the pytorch optimizer object optim_groups = [ { "params": [param_dict[pn] for pn in sorted(list(decay))], "weight_decay": float(config["weight_decay"]), }, { "params": [param_dict[pn] for pn in sorted(list(no_decay))], "weight_decay": 0.0, }, ] optimizer = torch.optim.AdamW( optim_groups, lr=lr, betas=(0.9, 0.99), eps=1e-7, ) lr_scheduler = LambdaLR( optimizer=optimizer, lr_lambda=lambda step: rate( step, model_size=model.module.params.dim, factor=lr_factor, min_factor=0.1, steps_in_epoch=num_batches * lr_decay_epochs, ), ) max_examples_per_rank = config.get("max_examples_per_rank", None) enable_checkpoint = bool(config["enable_checkpoint"]) checkpoint_path = config.get("checkpoint_path", "./checkpoints") enable_grad_scaler = bool(config["enable_grad_scaler"]) enable_shuffle = bool(config["enable_shuffle"]) enable_eval = bool(config["enable_eval"]) enable_cross_attention = bool(config["enable_cross_attention"]) Path(checkpoint_path).mkdir(parents=True, exist_ok=True) scaler = ( torch.cuda.amp.GradScaler(growth_interval=200) if enable_grad_scaler else None ) if rank == 0: wandb.init( project="composer-v18", config={ "batch_size": batch_size, "num_batches": num_batches, "accum_iter": accum_iter, "eval_iter": eval_iter, "lr_factor_times_world_size": lr_factor, "epochs": epochs, "layers": model.module.params.n_layers, "heads": model.module.params.n_heads, "dim": model.module.params.dim, "dropout": model.module.params.dropout, "label_smoothing": criterion.smoothing, "world_size": world_size, "max_examples_per_rank": max_examples_per_rank, "param_count": model.module.param_count, "enable_flash": model.module.params.enable_flash, "enable_grad_scaler": enable_grad_scaler, "enable_text_prompt": enable_text_prompt, "enable_eval": enable_eval, "enable_shuffle": enable_shuffle, "train_target": train_target, "enable_cross_attention": enable_cross_attention, }, ) wandb.watch(model, log=None) model.train() train_state = TrainState() for epoch in range(epochs): if rank == 0 and enable_shuffle: train_data_gen.shuffle() dist.barrier() _, _, _, train_state = run_epoch( train_data_gen.generate( force_batches=num_batches, order="random", seq_len_max=seq_len_max, ), lambda: eval_data_gen.generate( order="id", seq_len_max=seq_len_max, force_batches=100 ), # XXX model, SimpleLossCompute(criterion), optimizer, lr_scheduler, batch_size, mode="train", num_batches=num_batches, num_eval_batches=num_eval_batches, accum_iter=accum_iter, eval_iter=eval_iter if enable_eval else None, enable_checkpoint=enable_checkpoint and rank == 0, epoch=epoch, quiet=rank != 0, max_examples_per_rank=max_examples_per_rank, scaler=scaler, train_state=train_state, world_size=world_size, checkpoint_path=checkpoint_path, ) if __name__ == "__main__": import sys from util import configure_logging configure_logging() run(sys.argv[1]) wandb.finish()