from data import Vocab from model import ( Transformer, ModelArgs, SimpleLossCompute, LabelSmoothing, MusicalPositionEmbedTransformer, ) import torch from data_gen import Batch from embedder import Embedder import pytest import torch.nn as nn from typing import Optional class LittleTrainer: def __init__(self, model): super().__init__() self.model = model self.loss_compute = SimpleLossCompute(LabelSmoothing(0, 0.01)) self.optimizer = torch.optim.Adam( self.model.parameters(), lr=0.001, betas=(0.1, 0.15), eps=1e-7, ) self.scaler = torch.cuda.amp.GradScaler() def step(self, input, target=None, **kwargs): with torch.cuda.amp.autocast(dtype=torch.float16, enabled=True): self.optimizer.zero_grad(set_to_none=True) if isinstance(input, Batch): pred = self.model.forward( input.tgt, start_pos=0, **kwargs, ) loss_node = self.loss_compute(pred, input.tgt_y) else: pred = self.model.forward( input, start_pos=0, **kwargs, ) loss_node = self.loss_compute(pred, target) loss = loss_node.item() print(loss) self.scaler.scale(loss_node).backward() self.scaler.unscale_(self.optimizer) torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0) self.scaler.step(self.optimizer) self.scaler.update() del loss_node return loss, pred @pytest.fixture() def test_vocab(): return Vocab(1, 127, 1, 127, 20, 127, 2, 1, 127, 8, 24 * 4 * 8, 8, 24) @pytest.fixture() def untrained_model(test_vocab): smallargs = ModelArgs( 256, 2, 4, test_vocab.N, cache=True, enable_flash=False, max_seq_len=384, max_inference_seq_len=384, ) model = MusicalPositionEmbedTransformer(test_vocab, smallargs) model.eval() return model class TestEmbedder(nn.Module): def __init__(self, max_length=64, embed_dim=64): super().__init__() self.max_length = max_length self.embed_dim = embed_dim self.dummy_param = nn.Parameter(torch.empty(0)) def tokenize(self, sentences): device = next(self.parameters(recurse=True)).device z = torch.zeros( len(sentences), self.max_length, dtype=torch.long, device=device ) for i in range(len(sentences)): s = sentences[i].encode("utf-8") for j in range(min(len(s), self.max_length)): z[i, j] = s[j] return z, torch.ones(len(sentences), self.max_length, device=device) def forward( self, input_ids: torch.LongTensor, attention_mask: Optional[torch.FloatTensor] = None, ): assert input_ids.shape == attention_mask.shape return torch.normal( 0, 0.2, (input_ids.shape[0], input_ids.shape[1], self.embed_dim), device=input_ids.device, ) @pytest.fixture(params=[True, False]) def untrained_cross_model(test_vocab, request): cross_embed_dim = 64 smallargs = ModelArgs( 256, 4, 4, test_vocab.N, cache=True, enable_flash=request.param, enable_cross_attention=True, max_seq_len=384, max_inference_seq_len=384, cross_attention_embedding_dim=cross_embed_dim, ) model = MusicalPositionEmbedTransformer( test_vocab, smallargs, encoder=TestEmbedder(embed_dim=cross_embed_dim) ) model.eval() if request.param: model = model.cuda() return model @pytest.fixture(params=[True, False]) def untrained_embed_nocross_model(test_vocab, request): smallargs = ModelArgs( 256, 4, 4, test_vocab.N, cache=True, enable_flash=request.param, enable_cross_attention=False, max_seq_len=384, max_inference_seq_len=383, ) model = MusicalPositionEmbedTransformer( test_vocab, smallargs, encoder=TestEmbedder(embed_dim=256) ) model.eval() if request.param: model = model.cuda() return model