import os import re import math import time import json import boto3 import torch import random import tempfile import numpy as np import torch.nn as nn import torch.nn.functional as F from tqdm import tqdm from torch import autocast from typing import List from tokenizers import Tokenizer from contextlib import contextmanager # ----------------------------------------------------------------------------- # helpers # ----------------------------------------------------------------------------- def save_checkpoint(model, optimizer, global_step, checkpoint_dir, config): print(f"Saving checkpoint to {checkpoint_dir}") torch.save( { "model": model.state_dict(), "optimizer": optimizer.state_dict(), "global_step": global_step, "config": config, }, os.path.join(checkpoint_dir, f"last_ckpt.pth"), ) def load_model(model_filepath): """ Load a GPTModel from a checkpoint file. Args: model_filepath (str): Path to the model checkpoint file. Can be a local path or S3 path. Returns: tuple: (model, config) - The loaded model and its configuration """ with _download_from_s3_if_needed(model_filepath) as tmp_fp: checkpoint = torch.load(tmp_fp, map_location="cpu") config = checkpoint.get("config", GPTConfig()) model = GPTModel(config) # Load model state dict model.load_state_dict(checkpoint["model"]) return model, config def read_jsonl(filepath): data = [] with open(filepath) as f: for line in f: line = line.strip() if len(line) == 0: continue m = json.loads(line) data.append(m) return data def write_jsonl(data, filepath, do_append=False): openarg = "a" if do_append else "w" with open(filepath, openarg) as f: for d in data: f.write(json.dumps(d, ensure_ascii=False) + "\n") def get_filename(filepath, keep_ext=True): if "http" in filepath: clean_filepath = filepath.split("?")[0] else: clean_filepath = filepath filename = clean_filepath.split("/")[-1] if "." not in filename: raise ValueError("filename does not seem to contain a period.") m = re.search(r"(.+)\.([^\.]+)$", filename) if not m: raise ValueError(f"filename could not be parsed for `{filepath}`") filename = m.group(1) file_ext = m.group(2).lower() if len(file_ext) > 10: raise ValueError(f"file extension suspiciously long for `{filepath}`") if keep_ext: filename = filename + "." + file_ext return filename S3_BUCKET_PATH_RE = r"s3\:\/\/(.+?)\/" def _parse_s3_filepath(s3_filepath): bucket_name = re.search(S3_BUCKET_PATH_RE, s3_filepath).group(1) rel_s3_filepath = re.sub(S3_BUCKET_PATH_RE, "", s3_filepath) return bucket_name, rel_s3_filepath def download_s3_file( from_s3_filepath, to_local_filepath, ): bucket_name, from_rel_s3_filepath = _parse_s3_filepath(from_s3_filepath) client = boto3.client("s3") client.download_file(bucket_name, from_rel_s3_filepath, to_local_filepath) @contextmanager def _download_from_s3_if_needed(maybe_s3_filepath): tmp_filepath = maybe_s3_filepath if maybe_s3_filepath.startswith("s3://"): temp_dir = tempfile.TemporaryDirectory() filename = get_filename(maybe_s3_filepath, keep_ext=True) tmp_filepath = os.path.join(temp_dir.name, filename) download_s3_file(maybe_s3_filepath, tmp_filepath) yield tmp_filepath def load_tokenizer( tokenizer_filepath="s3://suno-data/georg/models/tokenizers/tokenizer_60k.json", ): with _download_from_s3_if_needed(tokenizer_filepath) as tmp_fp: tokenizer = Tokenizer.from_file(tmp_fp) tokenizer.add_special_tokens(["\n"]) tokenizer.pad_idx = tokenizer.token_to_id("[PAD]") return tokenizer # ----------------------------------------------------------------------------- # dataset of text + semantic tokens # -----------------------------------------------------------------------------`` MAX_TAG_LEN = 256 MAX_TOT_TAGS_LEN = 512 def _space_repl(m): s = m.group() n_newline = s.count("\n") if n_newline >= 2: return "\n\n" elif n_newline == 1: return "\n" return " " def _clean_tag(s, retain_newlines=False): s = s.replace("[", " ").replace("]", " ") return _simplify_whitespace(s, retain_newlines=retain_newlines) def _simplify_whitespace(text, retain_newlines=True): """simplify while respecting up to 2 newlines""" if retain_newlines: text = re.sub(r"\s+", _space_repl, text).strip() else: text = re.sub(r"\s+", " ", text).strip() return text def prepare_text_inference(tags: List[str], lyrics: str): tags = [ clean_tag for tag in tags if len(clean_tag := _clean_tag(tag, retain_newlines=False)) > 0 ] tags_str = ( f"{', '.join([tag[:MAX_TAG_LEN] for tag in tags])[:MAX_TOT_TAGS_LEN]}".strip() ) lyrics = _simplify_whitespace(lyrics, retain_newlines=True) # Combine tags and lyrics text = "" if len(tags_str) > 0: text += f"[{tags_str}]" if len(lyrics) > 0: if len(text) > 0: text += "\n\n" text += lyrics return text class SemanticDataset(torch.utils.data.Dataset): def __init__( self, dataset_dir: str, memmap_filename: str, metas_filename: str, semantic_n_tokens: int, cond_text_len: int = 2560, ): self.semantic_n_tokens = semantic_n_tokens self.cond_text_len = cond_text_len # open semantic memmap semantic_data = np.memmap( os.path.join(dataset_dir, memmap_filename), dtype=np.uint16, mode="r", ) semantic_data = semantic_data.reshape(-1, semantic_n_tokens, 1) self.semantic_data = semantic_data[:, :, 0] # load metas metas = read_jsonl(os.path.join(dataset_dir, metas_filename)) self.metas = metas assert len(self.semantic_data) == len(self.metas) print(f"Loaded {len(self.semantic_data)} samples") # load tokenizer self.tokenizer = load_tokenizer() def __len__(self): return self.semantic_data.shape[0] def __getitem__(self, idx): # semantic codes semantic_codes = torch.from_numpy(self.semantic_data[idx].copy()).long() # append eos token semantic_codes = torch.cat( [ torch.tensor([SEMANTIC_SOS_TOKEN]), semantic_codes, torch.tensor([SEMANTIC_EOS_TOKEN]), ] ) lyrics = self.metas[idx].get("text_aligned", self.metas[idx].get("text", "")) tags = self.metas[idx].get("tags", []) text = prepare_text_inference(tags, lyrics) # Build condition tensors text_codes = self.tokenizer.encode(text).ids[: self.cond_text_len] text_codes = text_codes + [self.tokenizer.pad_idx] * max( 0, self.cond_text_len - len(text_codes) ) # append infer token text_codes = text_codes text_codes = torch.tensor(text_codes).long() # for now assume all tokens are valid attention_mask = torch.ones(len(text_codes) + len(semantic_codes)).bool() return text_codes, semantic_codes, attention_mask def collate_fn(batch): text_input_ids, semantic_input_ids, attention_mask = zip(*batch) text_input_ids = torch.stack(text_input_ids) semantic_input_ids = torch.stack(semantic_input_ids) attention_mask = torch.stack(attention_mask) # Create input/target pairs for causal language modeling # the text will be pre-prompt and we only compute loss on the semantic tokens labels = semantic_input_ids[:, 1:] # get labels BEFORE modifying input_ids semantic_input_ids = semantic_input_ids[:, :-1] # then truncate input_ids attention_mask = attention_mask[:, :-1] # match input_ids length return { "text_input_ids": text_input_ids, "semantic_input_ids": semantic_input_ids, "attention_mask": attention_mask, "labels": labels, } # ----------------------------------------------------------------------------- # model # ----------------------------------------------------------------------------- class MultiHeadAttention(nn.Module): def __init__(self, config): super().__init__() assert ( config.hidden_size % config.num_heads == 0 ), "hidden_size must be divisible by num_heads" self.num_heads = config.num_heads self.hidden_size = config.hidden_size self.head_size = config.hidden_size // config.num_heads self.query = nn.Linear(config.hidden_size, config.hidden_size) self.key = nn.Linear(config.hidden_size, config.hidden_size) self.value = nn.Linear(config.hidden_size, config.hidden_size) self.proj = nn.Linear(config.hidden_size, config.hidden_size) self.dropout = nn.Dropout(config.dropout) def forward(self, x, attention_mask=None): B, T, C = x.size() # Split into heads q = self.query(x).view(B, T, self.num_heads, self.head_size).transpose(1, 2) k = self.key(x).view(B, T, self.num_heads, self.head_size).transpose(1, 2) v = self.value(x).view(B, T, self.num_heads, self.head_size).transpose(1, 2) # Scaled dot-product attention scale = math.sqrt(self.head_size) scores = torch.matmul(q, k.transpose(-2, -1)) / scale # Causal mask - prevent attending to future tokens causal_mask = torch.triu(torch.ones(T, T), diagonal=1).bool() scores.masked_fill_(causal_mask.to(scores.device), float("-inf")) # Apply attention mask if provided if attention_mask is not None: scores = scores.masked_fill( ~attention_mask.unsqueeze(1).unsqueeze(2), float("-inf") ) attn = F.softmax(scores, dim=-1) attn = self.dropout(attn) # Apply attention to values out = torch.matmul(attn, v) # Reshape and project out = out.transpose(1, 2).contiguous().view(B, T, C) out = self.proj(out) return out class TransformerBlock(nn.Module): def __init__(self, config): super().__init__() self.attention = MultiHeadAttention(config) self.mlp = nn.Sequential( nn.Linear(config.hidden_size, config.mlp_ratio * config.hidden_size), nn.GELU(), nn.Linear(config.mlp_ratio * config.hidden_size, config.hidden_size), nn.Dropout(config.dropout), ) self.ln1 = nn.LayerNorm(config.hidden_size) self.ln2 = nn.LayerNorm(config.hidden_size) self.dropout = nn.Dropout(config.dropout) def forward(self, x, attention_mask=None): # Pre-LayerNorm architecture x = x + self.dropout(self.attention(self.ln1(x), attention_mask)) x = x + self.dropout(self.mlp(self.ln2(x))) return x class GPTModel(nn.Module): def __init__(self, config): super().__init__() self.config = config self.token_embeddings = nn.Embedding(config.text_vocab_size, config.hidden_size) self.semantic_embeddings = nn.Embedding( config.semantic_vocab_size, config.hidden_size ) self.position_embeddings = nn.Embedding( config.max_position_embeddings, config.hidden_size ) self.dropout = nn.Dropout(config.dropout) self.blocks = nn.ModuleList( [TransformerBlock(config) for _ in range(config.num_layers)] ) self.ln_f = nn.LayerNorm(config.hidden_size) self.head = nn.Linear( config.hidden_size, config.semantic_vocab_size, bias=False ) # Initialize weights self.apply(self._init_weights) def _init_weights(self, module): if isinstance(module, nn.Linear): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) if module.bias is not None: torch.nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) elif isinstance(module, nn.LayerNorm): torch.nn.init.zeros_(module.bias) torch.nn.init.ones_(module.weight) def forward(self, text_input_ids, semantic_input_ids, attention_mask=None): B, T = text_input_ids.size() B, S = semantic_input_ids.size() # Get token embeddings token_emb = self.token_embeddings(text_input_ids) semantic_emb = self.semantic_embeddings(semantic_input_ids) # concatenate token and semantic embeddings input_emb = torch.cat([token_emb, semantic_emb], dim=1) # Add positional embeddings pos = torch.arange(0, T + S, dtype=torch.long, device=input_emb.device) pos_emb = self.position_embeddings(pos) x = self.dropout(input_emb + pos_emb) # Apply transformer blocks for block in self.blocks: x = block(x, attention_mask) x = self.ln_f(x) logits = self.head(x) return logits # Example configuration SEMANTIC_PAD_TOKEN = 4000 SEMANTIC_SOS_TOKEN = 4001 SEMANTIC_EOS_TOKEN = 4002 class GPTConfig: def __init__(self): self.text_vocab_size = 60004 self.semantic_vocab_size = 4003 self.max_position_embeddings = 751 + 751 self.hidden_size = 1536 self.num_layers = 12 self.num_heads = 12 self.mlp_ratio = 4 self.dropout = 0.1 self.semantic_pad_token = SEMANTIC_PAD_TOKEN self.semantic_sos_token = SEMANTIC_SOS_TOKEN self.semantic_eos_token = SEMANTIC_EOS_TOKEN def train_step(model, batch, optimizer): optimizer.zero_grad() text_input_ids = batch["text_input_ids"].cuda() semantic_input_ids = batch["semantic_input_ids"].cuda() attention_mask = batch["attention_mask"].cuda() labels = batch["labels"].cuda() # forward pass with autocast(device_type="cuda", dtype=torch.bfloat16): logits = model( text_input_ids=text_input_ids, semantic_input_ids=semantic_input_ids, attention_mask=attention_mask, ) # crop logits to the semantic tokens logits = logits[:, -semantic_input_ids.shape[1] :, :] # Reshape logits to [batch_size * sequence_length, vocab_size] logits = logits.reshape(-1, 4003) # Reshape labels to [batch_size * sequence_length] labels = labels.view(-1) loss = F.cross_entropy(logits, labels) loss.backward() optimizer.step() return loss if __name__ == "__main__": config = GPTConfig() model = GPTModel(config) # Convert model to bfloat16 model = model.to(dtype=torch.bfloat16) # model = torch.nn.DataParallel(model) # count parameters print(f"GPTModel: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M params") model.cuda() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4) # setup dataset train_dataset = SemanticDataset( dataset_dir="/app/suno/data/diffusion_mix/dac_vae_fixed_25hz", memmap_filename="data_semantic_val.bin", metas_filename="metas_val.jsonl", semantic_n_tokens=750, cond_text_len=750, ) train_dataloader = torch.utils.data.DataLoader( train_dataset, batch_size=20, shuffle=True, collate_fn=collate_fn, num_workers=8, ) # train global_step = 0 max_steps = 1_000_000 run_start_time = time.strftime("%Y-%m-%d_%H-%M-%S") checkpoint_dir = f"/app/suno/christian/checkpoints/gpt/{run_start_time}_s{random.randint(0, 9999)}" os.makedirs(checkpoint_dir, exist_ok=False) # start training while global_step < max_steps: pbar = tqdm(train_dataloader) for batch in pbar: loss = train_step(model, batch, optimizer) pbar.set_description(f"Loss: {loss.item()}") global_step += 1 if global_step % 100 == 0: save_checkpoint(model, optimizer, global_step, checkpoint_dir, config) if global_step >= max_steps: break