File size: 3,392 Bytes
30e9297
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
# Low-Level Design (LLD)

## Module Details

### s02_tokenizer.py β€” REMI Tokenizer

**Vocabulary Layout** (485 tokens total):
```
[0]        PAD
[1]        BOS (Beginning of Sequence)
[2]        EOS (End of Sequence)
[3]        SEP (Separator)
[4-131]    NoteOn (MIDI pitch 0-127)
[132-259]  NoteOff (MIDI pitch 0-127)
[260-291]  Velocity (32 bins, each = 4 MIDI velocity units)
[292-391]  TimeShift (10ms steps, 10ms-1000ms)
[392-451]  Tempo (40-200 BPM, 60 bins)
[452-483]  Position (32 positions per bar)
[484]      Bar (bar delimiter)
```

**Encoding Algorithm**:
1. Collect all non-drum notes across instruments
2. Sort by onset time, then pitch
3. For each note: emit [TimeShift, Position, Velocity, NoteOn, TimeShift(duration), NoteOff]
4. Insert Bar tokens at measure boundaries
5. Wrap with BOS/EOS

### s04_model.py β€” MusicTransformer

**Layer Stack** (per block):
```
Input β†’ RMSNorm β†’ GQA (with RoPE) β†’ Residual Add
      β†’ RMSNorm β†’ SwiGLU FFN       β†’ Residual Add β†’ Output
```

**Grouped Query Attention**:
- 8 query heads, 4 key-value heads
- Each KV head serves 2 query heads (n_rep = 2)
- Head dim = 256/8 = 32
- Uses PyTorch 2.0 `scaled_dot_product_attention` (Flash Attention backend when available)

**RoPE Implementation**:
- Precompute sin/cos frequencies: `freq[i] = 1 / (ΞΈ^(2i/d))`
- Apply rotation: `q' = q * cos + rotate_half(q) * sin`
- Device-compatible real-valued implementation (no complex tensors)

**SwiGLU FFN**:
- `output = W2(SiLU(W1(x)) * W3(x))`
- Hidden dim = 448 (nearest multiple of 64 to `2/3 * 4 * 256`)

### s05_trainer.py β€” Training Loop

**Optimization Details**:
- AdamW: β₁=0.9, Ξ²β‚‚=0.95, wd=0.1
- LR schedule: linear warmup (200 steps) β†’ cosine decay to 1e-6
- Gradient accumulation: 4 steps (effective batch = 4 Γ— 4 = 16)
- Max gradient norm: 1.0

**Memory Budget (estimated for 512 seq len)**:
| Component              | Memory     |
|------------------------|-----------|
| Model weights (FP32)   | ~20MB     |
| Gradients              | ~20MB     |
| Optimizer states        | ~40MB     |
| Activations (w/ ckpt)  | ~50MB     |
| Data batch             | ~4MB      |
| **Total**              | **~134MB**|

### s06_generator.py β€” Inference

**Decoding Pipeline**:
```
logits β†’ temperature_scale β†’ repetition_penalty β†’ top_k_filter β†’ top_p_filter β†’ softmax β†’ multinomial_sample
```

**KV-Cache**:
- Each layer stores (K, V) tensors after each forward pass
- New tokens only compute attention against cached KV + new KV
- Reset between generations to prevent cross-contamination

## Class Diagram

```
MusicTokenizer           MusicTransformer              Trainer
β”œβ”€β”€ midi_to_tokens()     β”œβ”€β”€ TransformerBlock (Γ—6)     β”œβ”€β”€ _train_epoch()
β”œβ”€β”€ tokens_to_midi()     β”‚   β”œβ”€β”€ GroupedQueryAttention  β”œβ”€β”€ _validate()
β”œβ”€β”€ decode_token()       β”‚   β”‚   β”œβ”€β”€ wq, wk, wv, wo    β”œβ”€β”€ _save_checkpoint()
β”œβ”€β”€ save()/load()        β”‚   β”‚   β”œβ”€β”€ RoPE application   └── load_checkpoint()
β”‚                        β”‚   β”‚   └── KV-cache
└── Token constants      β”‚   β”œβ”€β”€ SwiGLU FFN
                         β”‚   └── RMSNorm (Γ—2)
                         β”œβ”€β”€ token_emb (weight-tied)
                         β”œβ”€β”€ output projection
                         └── forward() / from_config()
```