import os import math import torch import einsum import numpy as np from torch import nn, einsum import torch.optim as optim from tqdm import tqdm from torch.nn.functional import mse_loss from suno_utils.utils.text import read_jsonl class CausalSelfAttention(nn.Module): def __init__(self, embed_size, num_heads): super().__init__() self.multihead_attn = nn.MultiheadAttention( embed_size, num_heads, batch_first=True ) def forward(self, x): seq_length = x.size(1) # Create a causal mask mask = torch.triu(torch.ones(seq_length, seq_length), diagonal=1).bool() mask = mask.to(x.device) # Apply causal self-attention attn_output, _ = self.multihead_attn(x, x, x, attn_mask=mask) return attn_output # 1. Model Architecture class GPTModel(nn.Module): def __init__( self, acoustic_codebook_size: int, n_acoustic_codebooks: int, embed_size: int, num_heads: int, num_layers: int, max_seq_length: int, ): super(GPTModel, self).__init__() self.token_embedding = nn.Embedding(acoustic_codebook_size, embed_size) self.position_embedding = nn.Embedding(max_seq_length, embed_size) self.layers = nn.ModuleList( [ nn.Sequential( CausalSelfAttention(embed_size, num_heads), nn.LayerNorm(embed_size), nn.Linear(embed_size, embed_size * 4), nn.GELU(), nn.Linear(embed_size * 4, embed_size), nn.LayerNorm(embed_size), ) for _ in range(num_layers) ] ) self.acoustic_heads = torch.nn.ModuleList() for n in range(n_acoustic_codebooks): self.acoustic_heads.append(nn.Linear(embed_size, acoustic_codebook_size)) def forward(self, x: torch.Tensor): bs, seq_length, n_codebooks = x.size() position_ids = torch.arange(seq_length, device=x.device).unsqueeze(0) token_embeds = self.token_embedding(x) position_embeds = self.position_embedding(position_ids) x = token_embeds + position_embeds # x: (bs, seq_length, n_codebooks, embed_size) # sum across codebook dimension x = x.sum(dim=-1) # (bs, seq_length, embed_size) for layer in self.layers: x = x + layer(x) # Residual connection # Apply acoustic heads outputs = [] for head in self.acoustic_heads: outputs.append(head(x)) # stack outputs along the last dimension outputs = torch.stack(outputs, dim=-1) return outputs def apply_delay_pattern(x, delay: int = 1, pad_token: int = 0): batch_size, seq_length, n_codebooks = x.size() # Calculate the maximum shift max_shift = delay * (n_codebooks - 1) # Create a new tensor filled with pad_token result = torch.full( (batch_size, seq_length + max_shift, n_codebooks), pad_token, dtype=x.dtype, device=x.device, ) for i in range(n_codebooks): shift = delay * i result[:, shift : shift + seq_length, i] = x[:, :, i] return result def restore_original_alignment(x, delay: int = 1, pad_token: int = 0): batch_size, extended_seq_length, n_codebooks = x.size() # Calculate the original sequence length original_seq_length = extended_seq_length - delay * (n_codebooks - 1) # Create a new tensor to store the result result = torch.zeros( batch_size, original_seq_length, n_codebooks, dtype=x.dtype, device=x.device ) for i in range(n_codebooks): shift = delay * i result[:, :, i] = x[:, shift : shift + original_seq_length, i] return result # 2. Dataset class MemmapDataset(torch.utils.data.Dataset): def __init__( self, acoustic_tokens_memmap_path: str, n_tokens_memmap: int, n_acoustic_codebooks: int, acoustic_codebook_size: int, ): acoustic_tokens = np.memmap( acoustic_tokens_memmap_path, dtype=np.uint16, mode="r" ) acoustic_tokens = acoustic_tokens.reshape( -1, n_tokens_memmap, n_acoustic_codebooks ) print(f"Acoustic tokens shape: {acoustic_tokens.shape}") self.acoustic_tokens = acoustic_tokens self.acoustic_codebook_size = acoustic_codebook_size self.pad_token = acoustic_codebook_size + 1 def __len__(self): return self.acoustic_tokens.shape[0] def __getitem__(self, idx): input_seq = torch.from_numpy(self.acoustic_tokens[idx, ...].copy()) # input_seq: (n_tokens_memmap, n_acoustic_codebooks) # Apply delay pattern to input sequence input_seq = apply_delay_pattern(input_seq, delay=1, pad_token=self.pad_token) # create target sequence by shifting input sequence by 1 target_seq = input_seq.clone() target_seq[:-1] = input_seq[1:] target_seq[-1] = self.pad_token return input_seq, target_seq # 25 hz codec # 10 sec chunks -> 250 codec time steps # 12 codebooks means 250 x 12 = 3000 tokens seq length for flat # delay pattern one should reduce that 250 + 1 = 251 tokens # 3. Training Loop def train(model, dataloader, optimizer, criterion, device): model.train() total_loss = 0 pbar = tqdm(dataloader, total=len(dataloader)) for input_seq, target_seq in pbar: input_seq, target_seq = input_seq.to(device), target_seq.to(device) optimizer.zero_grad() output = model(input_seq) loss = criterion(output.view(-1, output.size(-1)), target_seq.view(-1)) loss.backward() optimizer.step() total_loss += loss.item() pbar.set_description(f"Loss: {loss.item():.4f}") return total_loss / len(dataloader) if __name__ == "__main__": # Hyperparameters vocab_size = 1000 # Example value embed_size = 256 num_heads = 8 num_layers = 6 max_seq_length = 100 batch_size = 32 num_epochs = 10 learning_rate = 0.001 # Device configuration device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # Create model model = GPTModel(vocab_size, embed_size, num_heads, num_layers, max_seq_length).to( device ) # print number of model parameters in millions num_params = sum(p.numel() for p in model.parameters()) / 1_000_000 print(f"Number of GPT parameters: {num_params:.2f}M") dataset = MemmapDataset(coarse_memmap_path, max_seq_length) dataloader = torch.utils.data.DataLoader( dataset, batch_size=batch_size, shuffle=True ) # Loss and optimizer criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=learning_rate) # Training loop for epoch in range(num_epochs): loss = train(model, dataloader, optimizer, criterion, device)