Spaces:
Running
Running
Update core/music_generator.py
Browse files- core/music_generator.py +35 -2
core/music_generator.py
CHANGED
|
@@ -1,25 +1,58 @@
|
|
| 1 |
# core/music_generator.py
|
|
|
|
| 2 |
import torch
|
| 3 |
import torch.nn as nn
|
| 4 |
|
|
|
|
| 5 |
class MusicGenerator(nn.Module):
|
| 6 |
def __init__(self, embed_dim=128, nhead=4, num_layers=4, vocab_size=128):
|
| 7 |
super().__init__()
|
| 8 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 9 |
self.embed = nn.Embedding(vocab_size, embed_dim)
|
|
|
|
| 10 |
encoder_layer = nn.TransformerEncoderLayer(
|
| 11 |
d_model=embed_dim,
|
| 12 |
nhead=nhead,
|
| 13 |
batch_first=True
|
| 14 |
)
|
|
|
|
| 15 |
self.transformer = nn.TransformerEncoder(
|
| 16 |
encoder_layer,
|
| 17 |
num_layers=num_layers
|
| 18 |
)
|
|
|
|
| 19 |
self.fc = nn.Linear(embed_dim, vocab_size)
|
| 20 |
|
| 21 |
def forward(self, x):
|
|
|
|
|
|
|
|
|
|
| 22 |
x = self.embed(x)
|
| 23 |
x = self.transformer(x)
|
| 24 |
x = self.fc(x)
|
| 25 |
-
return x
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
# core/music_generator.py
|
| 2 |
+
|
| 3 |
import torch
|
| 4 |
import torch.nn as nn
|
| 5 |
|
| 6 |
+
|
| 7 |
class MusicGenerator(nn.Module):
|
| 8 |
def __init__(self, embed_dim=128, nhead=4, num_layers=4, vocab_size=128):
|
| 9 |
super().__init__()
|
| 10 |
+
|
| 11 |
+
# expose dimensions for external use
|
| 12 |
+
self.embedding_dim = embed_dim
|
| 13 |
+
self.vocab_size = vocab_size
|
| 14 |
+
|
| 15 |
+
# layers
|
| 16 |
self.embed = nn.Embedding(vocab_size, embed_dim)
|
| 17 |
+
|
| 18 |
encoder_layer = nn.TransformerEncoderLayer(
|
| 19 |
d_model=embed_dim,
|
| 20 |
nhead=nhead,
|
| 21 |
batch_first=True
|
| 22 |
)
|
| 23 |
+
|
| 24 |
self.transformer = nn.TransformerEncoder(
|
| 25 |
encoder_layer,
|
| 26 |
num_layers=num_layers
|
| 27 |
)
|
| 28 |
+
|
| 29 |
self.fc = nn.Linear(embed_dim, vocab_size)
|
| 30 |
|
| 31 |
def forward(self, x):
|
| 32 |
+
"""
|
| 33 |
+
x: (B, T) integer tokens
|
| 34 |
+
"""
|
| 35 |
x = self.embed(x)
|
| 36 |
x = self.transformer(x)
|
| 37 |
x = self.fc(x)
|
| 38 |
+
return x
|
| 39 |
+
|
| 40 |
+
def generate(self, seq_len=64, device="cpu"):
|
| 41 |
+
"""
|
| 42 |
+
Generate a random music token sequence
|
| 43 |
+
"""
|
| 44 |
+
self.eval()
|
| 45 |
+
|
| 46 |
+
# start with random token
|
| 47 |
+
tokens = torch.randint(0, self.vocab_size, (1, 1)).to(device)
|
| 48 |
+
|
| 49 |
+
for _ in range(seq_len - 1):
|
| 50 |
+
logits = self.forward(tokens)
|
| 51 |
+
next_token_logits = logits[:, -1, :]
|
| 52 |
+
|
| 53 |
+
probs = torch.softmax(next_token_logits, dim=-1)
|
| 54 |
+
next_token = torch.multinomial(probs, num_samples=1)
|
| 55 |
+
|
| 56 |
+
tokens = torch.cat([tokens, next_token], dim=1)
|
| 57 |
+
|
| 58 |
+
return tokens
|