# Model for ONNX conversion import torch import torch.nn as nn import torch.nn.functional as F from typing import Dict, Optional import math from model import ( ModelArgs, RMSNorm, ReZero, FeedForward, ) def precompute_freqs_cos_sins(dim: int, end: int, theta: float = 10000.0): freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)) t = torch.arange(end, device=freqs.device) freqs = torch.outer(t, freqs).float() freqs_cos = torch.cos(freqs) # real freqs_sin = torch.sin(freqs) # imag freqs_cos.requires_grad = False freqs_sin.requires_grad = False return freqs_cos, freqs_sin def apply_rotary_emb_real( x: torch.Tensor, freqs_cos: torch.Tensor, freqs_sin: torch.Tensor, ) -> torch.Tensor: seqlen, n_heads, _ = x.shape x_r, x_i = x.float().reshape(seqlen, n_heads, -1, 2).unbind(-1) _, _, head_dim = x_r.shape freqs_cos_v = freqs_cos.view(seqlen, 1, head_dim) freqs_sin_v = freqs_sin.view(seqlen, 1, head_dim) x_out_r = x_r * freqs_cos_v - x_i * freqs_sin_v x_out_i = x_r * freqs_sin_v + x_i * freqs_cos_v x_out = torch.stack((x_out_r, x_out_i), dim=-1).flatten(2) return x_out.type_as(x) def apply_rotary_emb_real_batched( x: torch.Tensor, freqs_cos: torch.Tensor, freqs_sin: torch.Tensor, ) -> torch.Tensor: bsz, seqlen, n_heads, _ = x.shape x_r, x_i = x.float().reshape(bsz, seqlen, n_heads, -1, 2).unbind(-1) _, _, _, head_dim = x_r.shape freqs_cos_v = freqs_cos.view(1, seqlen, 1, head_dim) freqs_sin_v = freqs_sin.view(1, seqlen, 1, head_dim) x_out_r = x_r * freqs_cos_v - x_i * freqs_sin_v x_out_i = x_r * freqs_sin_v + x_i * freqs_cos_v x_out = torch.stack((x_out_r, x_out_i), dim=-1).flatten(3) return x_out.type_as(x) class Attention(nn.Module): def __init__(self, args: ModelArgs): super().__init__() self.n_heads = args.n_heads self.head_dim = args.dim // args.n_heads self.wq = nn.Linear( args.dim, self.n_heads * self.head_dim, bias=False, dtype=torch.float32, ) self.wk = nn.Linear( args.dim, self.n_heads * self.head_dim, bias=False, dtype=torch.float32, ) self.wv = nn.Linear( args.dim, self.n_heads * self.head_dim, bias=False, dtype=torch.float32, ) self.wo = nn.Linear( self.n_heads * self.head_dim, args.dim, bias=False, dtype=torch.float32, ) self.cache = args.cache def forward( self, x: torch.Tensor, freqs_cos: torch.Tensor, freqs_sin: torch.Tensor, mask: torch.Tensor, ): bsz, seqlen, _ = x.shape xq, xk, xv = self.wq(x), self.wk(x), self.wv(x) xq = xq.view(bsz, seqlen, self.n_heads, self.head_dim) xk = xk.view(bsz, seqlen, self.n_heads, self.head_dim) xv = xv.view(bsz, seqlen, self.n_heads, self.head_dim) xq = apply_rotary_emb_real_batched( xq, freqs_cos=freqs_cos[:seqlen], freqs_sin=freqs_sin[:seqlen], ) xk = apply_rotary_emb_real_batched( xk, freqs_cos=freqs_cos[:seqlen], freqs_sin=freqs_sin[:seqlen], ) keys = xk values = xv xq = xq.transpose(1, 2) keys = torch.permute(keys, (0, 2, 3, 1)) values = values.transpose(1, 2) scores = torch.matmul(xq, keys) / math.sqrt(self.head_dim) scores = scores + mask scores = F.softmax(scores.float(), dim=-1).type_as(xq) output = torch.matmul(scores, values) # (bsz, n_local_heads, slen, head_dim) output = output.transpose(1, 2) return self.wo(output.contiguous().view(bsz, seqlen, -1)), xk, xv class AttentionOneStep(Attention): def __init__(self, args: ModelArgs): super().__init__(args) def forward( self, x: torch.Tensor, freqs_cos: torch.Tensor, freqs_sin: torch.Tensor, xk_in: torch.Tensor, xv_in: torch.Tensor, mask: torch.Tensor, ): seqlen, bsz, _, _ = xk_in.shape _, new_seqlen, _ = x.shape xk = xk_in.transpose(0, 1) xv = xv_in.transpose(0, 1) xq = self.wq(x).view(bsz, new_seqlen, self.n_heads, self.head_dim) xk_new = self.wk(x).view(bsz, new_seqlen, self.n_heads, self.head_dim) xv_new = self.wv(x).view(bsz, new_seqlen, self.n_heads, self.head_dim) xq = apply_rotary_emb_real_batched( xq, freqs_cos=freqs_cos[seqlen : seqlen + new_seqlen], freqs_sin=freqs_sin[seqlen : seqlen + new_seqlen], ) xk_new = apply_rotary_emb_real_batched( xk_new, freqs_cos=freqs_cos[seqlen : seqlen + new_seqlen], freqs_sin=freqs_sin[seqlen : seqlen + new_seqlen], ) xq = xq.transpose(1, 2) keys = torch.cat((xk, xk_new), dim=1) keys = torch.permute(keys, (0, 2, 3, 1)) values = torch.cat((xv, xv_new), dim=1).transpose(1, 2) scores = torch.matmul(xq, keys) / math.sqrt(self.head_dim) scores = scores + mask scores = F.softmax(scores.float(), dim=-1).type_as(xq) output = torch.matmul(scores, values) # (bsz, n_local_heads, slen, head_dim) output = output.transpose(1, 2) return ( self.wo(output.contiguous().view(bsz, new_seqlen, -1)), xk_new.transpose(0, 1), xv_new.transpose(0, 1), ) class CrossAttention(nn.Module): def __init__(self, args: ModelArgs): super().__init__() self.n_heads = args.n_heads self.head_dim = args.dim // args.n_heads self.embedding_dim = args.cross_attention_embedding_dim self.wq = nn.Linear( args.dim, self.n_heads * self.head_dim, bias=False, dtype=torch.float32, ) self.wk = nn.Linear( self.embedding_dim, self.n_heads * self.head_dim, bias=False, dtype=torch.float32, ) self.wv = nn.Linear( self.embedding_dim, self.n_heads * self.head_dim, bias=False, dtype=torch.float32, ) self.wo = nn.Linear( self.n_heads * self.head_dim, args.dim, bias=False, dtype=torch.float32, ) def forward( self, x: torch.Tensor, encoder_out: torch.Tensor, freqs_cos: torch.Tensor, freqs_sin: torch.Tensor, mask: torch.Tensor, ): bsz, seqlen, dim = x.shape ebsz, encoder_seqlen, edim = encoder_out.shape assert edim == self.embedding_dim assert bsz == ebsz xk, xv = self.wk(encoder_out), self.wv(encoder_out) xk = xk.view(ebsz, encoder_seqlen, self.n_heads, self.head_dim) xv = xv.view(ebsz, encoder_seqlen, self.n_heads, self.head_dim) xk = apply_rotary_emb_real_batched( xk, freqs_cos=freqs_cos[:encoder_seqlen], freqs_sin=freqs_sin[:encoder_seqlen], ) xq = self.wq(x) xq = xq.view(bsz, seqlen, self.n_heads, self.head_dim) xq = apply_rotary_emb_real_batched( xq, freqs_cos=freqs_cos[:seqlen], freqs_sin=freqs_sin[:seqlen], ) xq = xq.transpose(1, 2) keys = torch.permute(xk, (0, 2, 3, 1)).contiguous() values = xv.transpose(1, 2).contiguous() scores = torch.matmul(xq, keys) / math.sqrt(self.head_dim) scores = ( scores + mask ) # (ebsz, n_heads, seqlen, encoder_seqlen) + (ebsz, 1, seqlen, encoder_seqlen) scores = F.softmax(scores.float(), dim=-1).type_as(xq) output = torch.matmul(scores, values) # (bsz, n_heads, seqlen, head_dim) output = output.transpose(1, 2) # keys: (ebsz, n_heads, head_dim, encoder_seqlen) # values: (ebsz, n_heads, encoder_seqlen, head_dim) return self.wo(output.contiguous().view(bsz, seqlen, -1)), keys, values class CrossAttentionOneStep(CrossAttention): def __init__(self, args: ModelArgs): super().__init__(args) self.n_heads = args.n_heads self.head_dim = args.dim // args.n_heads def forward( self, x: torch.Tensor, start_pos: int, freqs_cos: torch.Tensor, freqs_sin: torch.Tensor, keys: torch.Tensor, values: torch.Tensor, mask: torch.Tensor, ): bsz, new_seqlen, _ = x.shape xq = self.wq(x) xq = xq.view(bsz, new_seqlen, self.n_heads, self.head_dim) xq = apply_rotary_emb_real_batched( xq, freqs_cos=freqs_cos[start_pos : start_pos + new_seqlen], freqs_sin=freqs_sin[start_pos : start_pos + new_seqlen], ) xq = xq.transpose(1, 2) scores = torch.matmul(xq, keys) / math.sqrt(self.head_dim) scores = ( scores + mask ) # (ebsz, n_heads, new_seqlen, encoder_seqlen) + (ebsz, 1, new_seqlen, encoder_seqlen) scores = F.softmax(scores.float(), dim=-1).type_as(xq) output = torch.matmul(scores, values) # (bsz, n_heads, new_seqlen, head_dim) output = output.transpose(1, 2) return self.wo(output.contiguous().view(bsz, new_seqlen, -1)) class TransformerBlock(nn.Module): def __init__(self, layer_id: int, args: ModelArgs): super().__init__() self.n_heads = args.n_heads self.dim = args.dim self.head_dim = args.dim // args.n_heads self.attention = Attention(args) self.feed_forward = FeedForward( dim=args.dim, hidden_dim=4 * args.dim, multiple_of=args.multiple_of, dropout_p=0.0, ) self.layer_id = layer_id self.attention_norm = RMSNorm(args.dim, eps=args.norm_eps) self.ffn_norm = RMSNorm(args.dim, eps=args.norm_eps) self.rezero = ReZero() self.enable_cross_attention = args.enable_cross_attention if args.enable_cross_attention: self.cross_attention = CrossAttention(args) self.cross_attention_norm = RMSNorm(args.dim, eps=args.norm_eps) else: self.cross_attention = None self.cross_attention_norm = None def forward( self, x: torch.Tensor, freqs_cos: torch.Tensor, freqs_sin: torch.Tensor, mask: torch.Tensor, encoder_out: Optional[torch.Tensor], encoder_mask: Optional[torch.Tensor], ): a_out, a_xk, a_xv = self.attention.forward( self.attention_norm(x), freqs_cos, freqs_sin, mask ) h = x + self.rezero(a_out) if self.cross_attention is not None: assert encoder_out is not None and encoder_mask is not None c_out, c_xk, c_xv = self.cross_attention.forward( self.cross_attention_norm(h), encoder_out, freqs_cos, freqs_sin, encoder_mask, ) h = h + self.rezero(c_out) else: c_xk, c_xv = None, None return ( h + self.rezero(self.feed_forward.forward(self.ffn_norm(h))), a_xk, a_xv, c_xk, c_xv, ) class TransformerBlockOneStep(TransformerBlock): def __init__(self, layer_id: int, args: ModelArgs): super().__init__(layer_id, args) self.attention = AttentionOneStep(args) if args.enable_cross_attention: self.cross_attention = CrossAttentionOneStep(args) def forward( self, x: torch.Tensor, a_xk: torch.Tensor, a_xv: torch.Tensor, c_xk: Optional[torch.Tensor], c_xv: Optional[torch.Tensor], freqs_cos: torch.Tensor, freqs_sin: torch.Tensor, mask: torch.Tensor, encoder_mask: Optional[torch.Tensor], ): start_pos, _, _, _ = a_xk.shape a_out, a_xk_new, a_xv_new = self.attention.forward( self.attention_norm(x), freqs_cos, freqs_sin, a_xk, a_xv, mask ) h = x + self.rezero(a_out) if self.cross_attention is not None: assert c_xk is not None and c_xv is not None and encoder_mask is not None c_out = self.cross_attention.forward( self.cross_attention_norm(h), start_pos, freqs_cos, freqs_sin, c_xk, c_xv, encoder_mask, ) h = h + self.rezero(c_out) return ( h + self.rezero(self.feed_forward.forward(self.ffn_norm(h))), a_xk_new, a_xv_new, ) class TransformerBlockSequence(nn.Module): def __init__(self, module_list): super().__init__() self.layers = module_list def forward( self, x: torch.Tensor, freqs_cos: torch.Tensor, freqs_sin: torch.Tensor, mask: torch.Tensor, encoder_out: Optional[torch.Tensor], encoder_mask: Optional[torch.Tensor], ): a_xks, a_xvs, c_xks, c_xvs = [], [], [], [] for layer in self.layers: x, a_xk, a_xv, c_xk, c_xv = layer( x, freqs_cos, freqs_sin, mask, encoder_out, encoder_mask ) a_xks.append(a_xk) a_xvs.append(a_xv) c_xks.append(c_xk) c_xvs.append(c_xv) a_xks_stack = torch.stack(a_xks, dim=2) a_xvs_stack = torch.stack(a_xvs, dim=2) if self.layers[0].enable_cross_attention: return ( x, a_xks_stack, a_xvs_stack, torch.stack(c_xks), torch.stack(c_xvs), ) else: return x, a_xks_stack, a_xvs_stack class TransformerBlockSequenceOneStep(nn.Module): def __init__(self, module_list): super().__init__() self.layers = module_list def forward( self, x: torch.Tensor, a_xks: torch.Tensor, a_xvs: torch.Tensor, c_xks: Optional[torch.Tensor], c_xvs: Optional[torch.Tensor], freqs_cos: torch.Tensor, freqs_sin: torch.Tensor, mask: torch.Tensor, encoder_mask: Optional[torch.Tensor], ): a_xk_news, a_xv_news = [], [] for i, layer in enumerate(self.layers): ( x, a_xk_new, a_xv_new, ) = layer( x, a_xks[:, :, i], a_xvs[:, :, i], c_xks[i] if c_xks is not None else None, c_xvs[i] if c_xvs is not None else None, freqs_cos, freqs_sin, mask, encoder_mask, ) a_xk_news.append(a_xk_new) a_xv_news.append(a_xv_new) return ( x, torch.stack(a_xk_news, dim=2), torch.stack(a_xv_news, dim=2), ) class Transformer(nn.Module): def __init__(self, params: ModelArgs): super().__init__() self.params = params self.vocab_size = params.vocab_size self.n_layers = params.n_layers self.tok_embeddings = nn.Embedding(params.vocab_size, params.dim) layers = torch.nn.ModuleList() for layer_id in range(params.n_layers): layers.append(TransformerBlock(layer_id, params)) self.layer_sequence = TransformerBlockSequence(layers) self.norm = RMSNorm(params.dim, eps=params.norm_eps) self.output = nn.Linear(params.dim, params.vocab_size, bias=False) self.freqs_cos, self.freqs_sin = precompute_freqs_cos_sins( self.params.dim // self.params.n_heads, self.params.max_inference_seq_len * 2, ) self.enable_cross_attention = params.enable_cross_attention def forward( self, tokens: torch.Tensor, aux_embeddings: torch.Tensor, encoder_out: Optional[ torch.Tensor ] = None, # expected (ebsz, encoder_seqlen, encoder_dim) encoder_valid: Optional[ torch.Tensor ] = None, # expected (ebsz, encoder_seqlen) of bools ): device = tokens.device h = self.tok_embeddings(tokens) self.freqs_cos = self.freqs_cos.to(device) self.freqs_sin = self.freqs_sin.to(device) h += aux_embeddings.to(device) seqlen = h.shape[1] mask = torch.full((seqlen, seqlen), float("-inf"), device=device) mask.triu_(diagonal=1) mask = mask.type_as(h) # encoder_mask should be (ebsz, 1, seqlen, encoder_seqlen) # where encoder_mask[..., i, j] = 0 if decoder token i can attend to encoder token j, else -inf if encoder_out is not None: assert encoder_valid is not None ebsz, encoder_seqlen = encoder_valid.shape assert ( encoder_out.shape[0] == ebsz and encoder_out.shape[1] == encoder_seqlen ) encoder_mask = torch.full( (ebsz, seqlen, encoder_seqlen), float("-inf"), device=device ) encoder_mask.masked_fill_(encoder_valid.unsqueeze(1), 0) encoder_mask = encoder_mask.unsqueeze(1) else: encoder_mask = None h, *cached_vals = self.layer_sequence( h, self.freqs_cos, self.freqs_sin, mask, encoder_out, encoder_mask, ) h = self.norm(h) # return the embeddings for the last non-pad token in each sequence # TODO update this idxs = ( torch.argmax((tokens == self.params.padding_idx).to(torch.int32), dim=-1) - 1 ) idxs[idxs == -1] = tokens.shape[1] - 1 embeddings = h[torch.arange(h.shape[0], dtype=torch.int32), idxs] return ( F.log_softmax(self.output(h), dim=-1), *cached_vals, embeddings, ) class TransformerOneStep(Transformer): def __init__(self, params: ModelArgs): super().__init__(params) layers = torch.nn.ModuleList() for layer_id in range(params.n_layers): layers.append(TransformerBlockOneStep(layer_id, params)) self.layer_sequence = TransformerBlockSequenceOneStep(layers) def forward( self, tokens: torch.Tensor, aux_embeddings: torch.Tensor, a_xks: torch.Tensor, a_xvs: torch.Tensor, c_xks: Optional[torch.Tensor], c_xvs: Optional[torch.Tensor], encoder_valid: Optional[ torch.Tensor ], # expected (ebsz, encoder_seqlen) of bools ): start_pos, _, _, _, _ = a_xks.shape seqlen = tokens.shape[1] h = self.tok_embeddings(tokens) self.freqs_cos = self.freqs_cos.to(h.device) self.freqs_sin = self.freqs_sin.to(h.device) h += aux_embeddings.to(h.device) mask = torch.full( (seqlen, start_pos + seqlen), float("-inf"), device=tokens.device ) mask.triu_(diagonal=1 + start_pos) mask = mask.type_as(h) # encoder_mask should be (ebsz, 1, new_seqlen, encoder_seqlen) # where encoder_mask[..., i, j] = 0 if decoder token i can attend to encoder token j, else -inf if c_xvs is not None: assert encoder_valid is not None ebsz, encoder_seqlen = encoder_valid.shape assert c_xvs.shape[1] == ebsz and c_xvs.shape[3] == encoder_seqlen encoder_mask = torch.full( (ebsz, seqlen, encoder_seqlen), float("-inf"), device=tokens.device ) encoder_mask.masked_fill_(encoder_valid.unsqueeze(1), 0) encoder_mask = encoder_mask.unsqueeze(1) else: encoder_mask = None h, a_xk_news, a_xv_news = self.layer_sequence( h, a_xks, a_xvs, c_xks, c_xvs, self.freqs_cos, self.freqs_sin, mask, encoder_mask, ) h = self.norm(h) return F.log_softmax(self.output(h), dim=-1), a_xk_news, a_xv_news class MusicalPositionEmbedTransformer(nn.Module): def __init__(self, vocab, params): super().__init__() # ONNX doesn't support flash attention assert not params.enable_flash self.vocab = vocab self.params = params self.transformer = Transformer(params) beats_max = vocab.embed_length_max // vocab.quantize_divisions self.bar_embedding = nn.Embedding( beats_max // 4 + 1, params.context_embedding_dim ) self.beat_embedding = nn.Embedding(4, params.context_embedding_dim) self.tick_embedding = nn.Embedding( vocab.quantize_divisions, params.context_embedding_dim ) self.polyphony_embedding = nn.Embedding( vocab.embed_polyphony_max + 1, params.context_embedding_dim ) self.context_mix = nn.Linear( params.context_embedding_dim * 8, params.dim, bias=False ) self.context_norm = RMSNorm(params.dim) self.param_count = 0 for p in self.parameters(): self.param_count += p.numel() self.dim = params.dim def get_context(self, x): context = self.context_mix( torch.cat( [ self.bar_embedding(x[..., 1]), self.beat_embedding(x[..., 2]), self.tick_embedding(x[..., 3]), self.polyphony_embedding(x[..., 4]), self.bar_embedding(x[..., 5]), self.beat_embedding(x[..., 6]), self.tick_embedding(x[..., 7]), self.polyphony_embedding(x[..., 8]), ], dim=-1, ) ) return self.context_norm(context) def forward( self, x: torch.Tensor, encoder_out: Optional[torch.Tensor] = None, encoder_valid: Optional[torch.Tensor] = None, ): return self.transformer( x[..., 0], aux_embeddings=self.get_context(x), encoder_out=encoder_out, encoder_valid=encoder_valid, ) class MusicalPositionEmbedTransformerOneStep(MusicalPositionEmbedTransformer): def __init__(self, vocab, params): super().__init__(vocab, params) self.transformer = TransformerOneStep(params) def forward( self, x: torch.Tensor, a_xks: torch.Tensor, a_xvs: torch.Tensor, c_xks: Optional[torch.Tensor] = None, c_xvs: Optional[torch.Tensor] = None, encoder_valid: Optional[torch.Tensor] = None, ): return self.transformer( x[..., 0], aux_embeddings=self.get_context(x), a_xks=a_xks, a_xvs=a_xvs, c_xks=c_xks, c_xvs=c_xvs, encoder_valid=encoder_valid, ) class MusicalPositionEmbedTransformerOneStepNoCross( MusicalPositionEmbedTransformerOneStep ): def __init__(self, vocab, params): super().__init__(vocab, params) def forward( self, x: torch.Tensor, a_xks: torch.Tensor, a_xvs: torch.Tensor, ): return super().forward(x, a_xks, a_xvs)