Musombi commited on
Commit
d1d27cf
·
verified ·
1 Parent(s): 9350d53

Update core/music_generator.py

Browse files
Files changed (1) hide show
  1. 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
- self.embedding_dim = embed_dim # <--- ADD THIS
 
 
 
 
 
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