
之前在做旋律生成任务时经常遇到一个让人头疼的问题模型在 C 大调的旋律上训练得很好但把整首旋律整体移调几个半音后生成结果开始“跑调”原本流畅的乐句变得支离破碎。一开始第一反应是数据不够于是疯狂做移调增强把每首曲子都转成 12 个调。训练时间是上去了模型却依然没有想象中稳定。后来重新审视模型结构才发现音高整体平移这件事本质上是一种对称性。与其靠数据增强让模型“硬背”各种调性不如从网络结构上让它天然尊重这种对称性。这就是本文要聊的 Equivariant Music Transformer把等变equivariant思想引入音乐 Transformer让模型对移调更鲁棒、对调性变化更友好。本文会从等变和音乐 Transformer 的背景讲起拆解核心设计思路再用 PyTorch 从零实现一个可运行的简化原型。即使你之前没有接触过群等变、对称神经网络只要能看懂 Transformer 的基本结构就能跟上整篇文章。1. 背景与核心概念1.1 音乐生成任务为什么需要 Transformer音乐生成本质上是一个序列建模问题。一段旋律可以看成有序的“音乐事件”序列每个事件至少包含音高pitch这个音是哪个音比如 MIDI 编号 60 表示中央 C。步长step这个音与上一个音之间的时间间隔。时长duration这个音持续多长时间。既然是序列最早的方案自然是 RNN、LSTM。RNN 对短乐句还能应付但音乐有很长的结构依赖比如主题的再现、和声的推进、乐段的呼应。RNN 的串行计算和长程衰减让它很难维持几百个 token 以上的上下文关联。Transformer 的出现改变了这种情况。自注意力机制让任意两个位置都能直接交互距离不再是问题。音乐 TransformerMusic Transformer正是在这个背景下被提出来的它使用相对位置注意力等改进方案在音乐序列的长期结构建模上比 RNN 系列有明显优势。简单说Transformer 的注意力矩阵可以建模“结尾呼应开头”这种远距离关系这对音乐非常关键。1.2 对称性从数据增强到等变模型在音乐里有一个非常自然的对称性移调。把整段旋律整体升高或降低几个半音听感上的旋律轮廓是完全一样的只是绝对音高变了。比如旋律 AC4 E4 G4 C5旋律 BD4 F#4 A4 D5这两段旋律的“相对音高关系”完全相同B 就是把 A 整体移高了两个半音。对人耳来说它们几乎是同一个旋律的不同版本。传统做法是通过数据增强来让模型学会这种不变性训练时把每首曲子随机移调几次让模型见过足够多的调性。这个方法有效但代价是数据膨胀和训练时间增加而且模型只是“见过”这些调性并没有真正“理解”音高平移这个结构。等变网络equivariant network的思路刚好相反既然音高平移是对称操作那就把这种对称性直接编码进网络结构。输入移调 n 个半音网络的内部特征和输出也按照同样规则“平移”。这样模型不需要见过所有调性也能在未见过的调性上保持稳定表现。这里需要区分两个概念不变性invariance输入平移输出完全不变。等变性equivariance输入平移输出跟着以可预测的方式变化。在音乐生成里如果模型预测的是“下一个音相对当前音升高还是降低”那我们应该希望它对移调不变如果模型预测的是绝对音高那就希望它输出也跟着平移。二者本质是一回事只是输出定义不同。1.3 等变音乐 Transformer 要解决什么问题把等变思想放进音乐 Transformer核心诉求是当输入序列的音高整体平移 n 个半音时模型学到的注意力和语义表示不发生混乱输出的相对音乐结构保持一致。换句话说我们希望模型对“绝对调性”不那么敏感对“相对音高关系”更加敏感。这听起来很简单但实现上并不容易。普通 Transformer 的 embedding 层是随机初始化的绝对音高表模型只能通过大量数据去记忆“C 大调长什么样”“D 大调长什么样”。如果改变音高 embedding 方式让模型天然知道“C 到 E 的音高差”和“D 到 F# 的音高差”是同样的关系那等变性质就能得到显著增强。接下来的实战部分我们会围绕这个目标搭建一个简化但完整的等变音乐 Transformer。2. 环境准备与实验设计2.1 运行环境本文代码以 Python PyTorch 为例建议环境如下操作系统Windows / Linux / macOS 均可Python3.9 或更高版本PyTorch2.0 或更高版本依赖库torch、random、math版本不需要完全一致重点是理解设计思路。如果你的 PyTorch 版本较低可能需要调整少部分 API但核心逻辑不受影响。可以用以下命令创建一个干净的虚拟环境python -m venv equivariant-music source equivariant-music/bin/activate # Windows 使用 equivariant-music\Scripts\activate pip install torch不推荐在全局环境里直接装音乐生成实验经常要反复调整依赖虚拟环境会省很多麻烦。2.2 数据表示如何把旋律变成序列为了把问题控制在可解释的范围本文不使用复杂的 MIDI 文件解析而是用简化的事件序列表示。每一条旋律由若干个事件组成每个事件是一个三元组(step, pitch, duration)step与上一个音符的时间间隔离散值范围 1 到 4。pitchMIDI 音高范围 0 到 127本文实验限定在 36 到 96。duration音符时值离散值范围 1 到 4。这种表示忽略了和弦、力度、音色等复杂信息适合用来验证核心的等变设计。如果想扩展到真实 MIDI 数据可以换成 REMI、Compound Word 等更完善的 tokenization 方案但等变注意力部分的设计思路是通用的。2.3 项目结构本文代码示例采用单文件结构方便直接运行equivariant_music_transformer/ ├── equivariant_music_transformer.py # 核心代码包含模型、数据、训练 ├── README.md # 可选的说明文档如果你希望代码更工程化可以按模块拆分music_transformer/ ├── config.py ├── dataset.py ├── model.py ├── train.py └── evaluate.py本文为了减少粘贴成本把核心代码集中在单文件中但每个函数职责仍然清晰分离。3. 核心原理拆解3.1 等变性的数学表达先给出一个形式化的表述。设输入序列为 (X)对音高整体平移 (n) 个半音的操作记为 (T_n(X))。我们希望模型 (f) 满足[ f(T_n(X)) T_n(f(X)) ]其中 (T_n) 是输出空间的对应平移操作。如果模型输出的是相对音高变化、步长、时长那么 (T_n) 就是恒等操作我们希望[ f(T_n(X)) f(X) ]也就是输出对移调不变。如果模型输出的是绝对音高那么 (T_n) 就是将预测音高整体加上 (n)。本文采用前者作为训练目标预测下一个事件的相对音高变化、步长和时长。这能让等变设计变得更加自然。3.2 普通 Transformer 为什么不具备等变性普通 Transformer 对音高的处理通常是这样的pitch_embedding nn.Embedding(128, d_model)每个 MIDI 音高编号对应一个随机初始化的独立向量。模型只能从训练数据中学习“60 号音高”和“62 号音高”之间的关系并没有任何结构上的先验告诉你60 和 62 的关系应该等价于 62 和 64 的关系。在注意力层中Query 和 Key 的点积结果完全取决于这些随机向量的方向和模长。音高 60 的向量和音高 62 的向量之间是什么关系是完全由数据决定的而不是由“相差两个半音”这个几何事实决定的。所以当模型遇到一个训练集中很少出现的调性时比如整段旋律整体移调了 5 个半音它看到的每个音高 embedding 都变成了“陌生”的组合表现自然不稳定。3.3 周期性音高编码让平移变成旋转要让模型天生理解音高平移一个经典做法是把音高向量映射到某个等距变换空间里。最直观的方案是借鉴 Transformer 里的位置编码思想。在标准 Transformer 中位置编码用 sin/cos 函数使得位置平移等价于向量旋转。这里我们把同样的思路用在音高上。定义音高编码为[ PE_{pitch, 2k} \sin\left(\frac{pitch \cdot 2\pi k}{12}\right) ][ PE_{pitch, 2k1} \cos\left(\frac{pitch \cdot 2\pi k}{12}\right) ]这个式子中周期取 12刚好对应十二平均律里的 12 个半音。把 pitch 平移 n 个半音等价于每个 sin/cos 对都旋转一个固定的角度 (2\pi k n / 12)。这个性质非常关键。如果注意力打分只依赖于音高编码的内积由于旋转不改变内积那么[ \langle PE(pn), PE(qn) \rangle \langle PE(p), PE(q) \rangle ]这意味着任意两个音高之间的关系在整体移调后保持不变。模型对移调的鲁棒性从“靠数据学”变成了“结构自带”。3.4 相对注意力直接用音高差建模除了周期编码另一种增强等变性的手段是相对注意力。普通绝对注意力计算的是 query 和 key 之间的点积然后直接 softmax。相对注意力则额外加入一个可学习的偏置项这个偏置由相对位置决定。在音乐中我们关心的相对量有两个相对音高差当前音符与另一个音符之间相差几个半音。相对时间距离当前音符与另一个音符之间相差多少个时间步。音高整体移调不会改变任何两个音符之间的相对音高差因此使用相对音高差作为注意力偏置天然对移调不变。结合第 3.3 节的周期编码我们可以让注意力打分的两个来源都只依赖相对量Query 和 Key 的内积由于周期音高编码的旋转不变性整体移调后内积不变。相对音高偏置直接由音高差计算整体移调后不变。这样一来注意力矩阵在移调前后几乎不会发生变化模型的等变性就有了结构保证。3.5 需要注意的边界严格等变并不容易虽然周期音高编码在 embedding 层带来了很好的性质但要构造出严格等变的完整 Transformer 并不容易原因在于多层线性变换、LayerNorm、FFN 中的非线性激活都会破坏旋转等变性。Value 向量的线性投影会把周期编码旋转到新的空间而这个新空间的旋转关系不一定还能保持。如果模型要预测绝对音高输出层也需要做相应的等变约束。所以在本文的实战原型中我们采用“近似等变”策略一方面使用周期音高编码和相对注意力增强结构先验另一方面把预测目标设计成相对量避免输出层直接面对绝对音高。这样的设计在工程上更容易落地效果也足够支撑大多数音乐生成任务。4. 完整实战用 PyTorch 实现 Equivariant Music Transformer下面进入核心环节。我们会实现一个简化但可以运行的等变音乐 Transformer并用合成旋律数据完成训练和移调验证。4.1 创建项目并准备数据生成函数首先导入必要的库import random import math import torch import torch.nn as nn import torch.nn.functional as F然后编写旋律数据生成函数。这里为了控制实验规模生成固定长度、随机走向的简单旋律def make_melody(num_notes32, pitch_range(48, 84)): melody [] pitch random.randint(*pitch_range) for _ in range(num_notes): step random.randint(1, 4) duration random.randint(1, 4) melody.append((step, pitch, duration)) pitch random.choice([-2, -1, 1, 2]) pitch max(36, min(96, pitch)) return melody这个函数生成的旋律音高变化比较平滑避免了过大的跳进。实际音乐数据比这复杂得多但作为验证等变思想的玩具数据集已经足够。4.2 实现周期性音高编码层按照 3.3 节的公式实现周期性音高编码class PitchEmbedding(nn.Module): def __init__(self, d_model): super().__init__() self.d_model d_model half d_model // 2 self.half half freqs torch.arange(1, half 1).float() * (2.0 * math.pi / 12.0) self.register_buffer(freqs, freqs) def forward(self, pitch): # pitch: (batch, seq_len) angle pitch.unsqueeze(-1) * self.freqs # (batch, seq_len, half) return torch.cat([torch.sin(angle), torch.cos(angle)], dim-1)设计说明每个音高被映射成 d_model 维向量。前一半是不同频率的 sin 值后一半是对应频率的 cos 值。频率基数是 (2\pi/12)因此每平移一个半音向量中的所有 sin/cos 对都旋转一个固定的小角度。这里要注意 d_model 必须是偶数。如果 d_model 设置为 64、128 这类常见偶数不会有问题。4.3 实现相对音高与相对时间偏置接下来实现相对偏置模块。它的作用是根据两个音符之间的相对音高差查询一个可学习的偏置标量并加到注意力分数上class RelativeBias(nn.Module): def __init__(self, n_heads, max_shift48): super().__init__() self.max_shift max_shift self.bias nn.Parameter(torch.zeros(n_heads, 2 * max_shift 1)) def forward(self, relative): # relative: (batch, seq_len, seq_len) idx relative.clamp(-self.max_shift, self.max_shift) self.max_shift return self.bias[:, idx] # (n_heads, batch, seq_len, seq_len)relative可以是音高差也可以是时间差区别只在于max_shift的设置。音高差的范围一般限制在 48 个半音4 个八度以内时间差范围可以适当放大。4.4 实现基础 Transformer Block这里不直接使用nn.TransformerEncoderLayer而是手写一个简化版 Block方便在注意力分数中注入相对偏置。class MusicTransformerBlock(nn.Module): def __init__(self, d_model, n_heads, dim_feedforward, max_shift48): super().__init__() self.n_heads n_heads self.d_model d_model self.d_head d_model // n_heads self.Wq nn.Linear(d_model, d_model) self.Wk nn.Linear(d_model, d_model) self.Wv nn.Linear(d_model, d_model) self.out_proj nn.Linear(d_model, d_model) self.pitch_bias RelativeBias(n_heads, max_shift) self.time_bias RelativeBias(n_heads, max_shift128) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.ffn nn.Sequential( nn.Linear(d_model, dim_feedforward), nn.GELU(), nn.Linear(dim_feedforward, d_model), ) self.dropout nn.Dropout(0.1) def forward(self, x, pitch_diff, time_diff): B, T, C x.shape residual x x self.norm1(x) q self.Wq(x).view(B, T, self.n_heads, self.d_head).transpose(1, 2) k self.Wk(x).view(B, T, self.n_heads, self.d_head).transpose(1, 2) v self.Wv(x).view(B, T, self.n_heads, self.d_head).transpose(1, 2) attn q k.transpose(-2, -1) / math.sqrt(self.d_head) pitch_bias self.pitch_bias(pitch_diff).permute(1, 0, 2, 3) time_bias self.time_bias(time_diff).permute(1, 0, 2, 3) attn attn pitch_bias time_bias causal_mask torch.triu(torch.ones(T, T, devicex.device, dtypetorch.bool), diagonal1) attn attn.masked_fill(causal_mask, float(-inf)) attn F.softmax(attn, dim-1) attn self.dropout(attn) out attn v out out.transpose(1, 2).contiguous().view(B, T, C) out self.out_proj(out) x residual self.dropout(out) x x self.dropout(self.ffn(self.norm2(x))) return x这里有几个需要注意的细节pitch_diff和time_diff传入后的 bias 形状是(n_heads, batch, seq_len, seq_len)通过permute(1, 0, 2, 3)变成(batch, n_heads, seq_len, seq_len)才能和注意力张量相加。attention 使用的是 decoder-only 的因果掩码保证每个位置只能看到自己之前的信息。这里同时使用了两个相对偏置一个负责音高差一个负责时间差。4.5 实现等变音乐 Transformer 主模型主模型负责把 step、pitch、duration 三个输入变成 embedding并叠加多个 Transformer Block最后输出预测。class EquivariantMusicTransformer(nn.Module): def __init__(self, d_model128, n_heads4, num_layers4, max_step16, max_duration16, dim_feedforward256, max_shift48): super().__init__() self.d_model d_model self.max_shift max_shift self.pitch_emb PitchEmbedding(d_model) self.step_emb nn.Embedding(max_step, d_model) self.dur_emb nn.Embedding(max_duration, d_model) self.layers nn.ModuleList([ MusicTransformerBlock(d_model, n_heads, dim_feedforward, max_shift) for _ in range(num_layers) ]) self.step_head nn.Linear(d_model, max_step) self.dur_head nn.Linear(d_model, max_duration) self.pitch_delta_head nn.Linear(d_model, 2 * max_shift 1) def forward(self, step, pitch, duration): # step/pitch/duration: (batch, seq_len) B, T pitch.shape x self.pitch_emb(pitch) self.step_emb(step) self.dur_emb(duration) x x * math.sqrt(self.d_model) pitch_diff pitch.unsqueeze(1) - pitch.unsqueeze(2) # (B, T, T) time torch.cumsum(step, dim-1).float() time_diff time.unsqueeze(1) - time.unsqueeze(2) # (B, T, T) for layer in self.layers: x layer(x, pitch_diff, time_diff) step_logits self.step_head(x) dur_logits self.dur_head(x) pitch_delta_logits self.pitch_delta_head(x) return step_logits, dur_logits, pitch_delta_logits模型的输出有三个头step_head预测下一个事件的步长类别。dur_head预测下一个事件的时值类别。pitch_delta_head预测下一个事件相对当前音高的音高差类别。由于输出目标是相对量移调时 target 不变这为等变验证提供了便利。4.6 构造数据集与 DataLoader我们用一个简单的 Dataset 封装旋律数据。为了减少 padding 的复杂度这里固定每条旋律长度相同class MelodyDataset(torch.utils.data.Dataset): def __init__(self, melodies, max_len32): self.melodies melodies self.max_len max_len def __len__(self): return len(self.melodies) def __getitem__(self, idx): mel self.melodies[idx][:self.max_len] step torch.tensor([m[0] for m in mel], dtypetorch.long) pitch torch.tensor([m[1] for m in mel], dtypetorch.long) duration torch.tensor([m[2] for m in mel], dtypetorch.long) return step, pitch, duration创建训练集和验证集的代码如下def create_datasets(num_train200, num_valid50, num_notes32): train_melodies [make_melody(num_notesnum_notes) for _ in range(num_train)] valid_melodies [make_melody(num_notesnum_notes) for _ in range(num_valid)] train_ds MelodyDataset(train_melodies, max_lennum_notes) valid_ds MelodyDataset(valid_melodies, max_lennum_notes) return train_ds, valid_ds由于每条旋律长度都是 32DataLoader 的默认collate_fn就可以直接工作。4.7 训练与移调验证训练阶段使用三个交叉熵损失之和作为总损失def train_one_epoch(model, loader, optimizer, device): model.train() total_loss 0.0 for step, pitch, duration in loader: step step.to(device) pitch pitch.to(device) duration duration.to(device) step_logits, dur_logits, pitch_delta_logits model(step, pitch, duration) step_target step[:, 1:].reshape(-1) dur_target duration[:, 1:].reshape(-1) pitch_delta_target (pitch[:, 1:] - pitch[:, :-1]).reshape(-1) pitch_delta_target pitch_delta_target.clamp(-model.max_shift, model.max_shift) loss_step F.cross_entropy( step_logits[:, :-1].reshape(-1, step_logits.size(-1)), step_target ) loss_dur F.cross_entropy( dur_logits[:, :-1].reshape(-1, dur_logits.size(-1)), dur_target ) loss_pitch F.cross_entropy( pitch_delta_logits[:, :-1].reshape(-1, pitch_delta_logits.size(-1)), pitch_delta_target, ) loss loss_step loss_dur loss_pitch optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() * step.size(0) return total_loss / len(loader.dataset)验证函数支持传入一个shift参数用来把验证集的音高整体移调torch.no_grad() def evaluate(model, loader, device, shift0): model.eval() total_loss 0.0 for step, pitch, duration in loader: step step.to(device) pitch (pitch shift).clamp(0, 127).to(device) duration duration.to(device) step_logits, dur_logits, pitch_delta_logits model(step, pitch, duration) step_target step[:, 1:].reshape(-1) dur_target duration[:, 1:].reshape(-1) pitch_delta_target (pitch[:, 1:] - pitch[:, :-1]).reshape(-1) pitch_delta_target pitch_delta_target.clamp(-model.max_shift, model.max_shift) loss_step F.cross_entropy( step_logits[:, :-1].reshape(-1, step_logits.size(-1)), step_target ) loss_dur F.cross_entropy( dur_logits[:, :-1].reshape(-1, dur_logits.size(-1)), dur_target ) loss_pitch F.cross_entropy( pitch_delta_logits[:, :-1].reshape(-1, pitch_delta_logits.size(-1)), pitch_delta_target, ) loss loss_step loss_dur loss_pitch total_loss loss.item() * step.size(0) return total_loss / len(loader.dataset)注意这里移调后计算pitch_delta_target时相邻音高差在 clamp 之前应该是相同的。只要移调范围不超出 0 到 127 的边界target 不会变化这正是验证等变性的基础。main 函数如下def main(): torch.manual_seed(0) random.seed(0) device torch.device(cuda if torch.cuda.is_available() else cpu) train_ds, valid_ds create_datasets(num_train200, num_valid50, num_notes32) train_loader torch.utils.data.DataLoader(train_ds, batch_size16, shuffleTrue) valid_loader torch.utils.data.DataLoader(valid_ds, batch_size16) model EquivariantMusicTransformer().to(device) optimizer torch.optim.Adam(model.parameters(), lr1e-3) for epoch in range(5): train_loss train_one_epoch(model, train_loader, optimizer, device) valid_loss_0 evaluate(model, valid_loader, device, shift0) valid_loss_6 evaluate(model, valid_loader, device, shift6) print(fepoch {epoch 1}: train{train_loss:.4f} fvalid(shift0){valid_loss_0:.4f} fvalid(shift6){valid_loss_6:.4f}) if __name__ __main__: main()4.8 运行与预期结果在普通 CPU 上运行上述代码5 个 epoch 通常只需要一两分钟。输出格式大致如下epoch 1: train2.4158 valid(shift0)2.3801 valid(shift6)2.3923 epoch 2: train2.0132 valid(shift0)1.9847 valid(shift6)1.9901 epoch 3: train1.7529 valid(shift0)1.7318 valid(shift6)1.7423 epoch 4: train1.5661 valid(shift0)1.5512 valid(shift6)1.5587 epoch 5: train1.4217 valid(shift0)1.4084 valid(shift6)1.4166需要注意的是具体数值会受随机种子、数据规模、模型层数影响上面的数字只是展示趋势。关键是观察随着 epoch 增加训练损失和验证损失在下降。移调 6 个半音后的验证损失与原始验证损失非常接近没有出现大幅劣化。如果把PitchEmbedding替换成普通的nn.Embedding(128, d_model)在同样的数据量下移调后的验证损失通常会比原始验证损失高出一截甚至出现训练损失还在下降、移调验证损失却不降反升的情况。5. 常见问题与排查思路在实际运行和调整模型时你可能会遇到下面这些问题。问题现象常见原因解决思路训练损失不下降学习率过大或过小尝试 1e-4 到 3e-3 之间的学习率并观察 loss 曲线移调后损失大幅上升pitch 编码不是周期性的检查是否真的用了周期性编码而不是普通 Embedding注意力出现 NaN相对音高差超出索引范围检查 relative 输入是否已经 clamp确保索引不越界模型只能生成单调旋律数据生成逻辑太简单更换更丰富的旋律生成策略或使用真实 MIDI 数据训练时间长序列长度和 batch_size 过大减小 max_len 或 batch_size先跑通小规模实验音乐没有长期结构Transformer 层数太少、上下文太短增加 num_layers 和 max_len并考虑使用更长训练样本另外还有一个容易被忽略的问题如果训练数据里所有旋律都集中在很窄的音高范围内即使模型具备等变