基于Keras构建非图像生成GAN示例:音乐生成实现咨询
Keras实现序列数据GAN(适配音乐生成场景)
音乐生成属于序列数据生成任务,和图像的二维网格结构不同,它依赖时序上下文,因此GAN的网络结构需要改用循环神经网络(LSTM/GRU)来处理序列依赖。以下是一个可直接运行的Keras示例,你可以基于此适配音乐数据(比如MIDI转成整数序列)。
1. 准备依赖与模拟序列数据
音乐数据通常会被预处理为固定长度的整数序列(每个整数对应一个音符/和弦),这里先模拟一批训练数据:
import numpy as np from tensorflow.keras.models import Model from tensorflow.keras.layers import Input, LSTM, Dense, Reshape from tensorflow.keras.optimizers import Adam from tensorflow.keras.utils import to_categorical # 模拟参数:序列长度、词汇表大小(对应不同音符)、样本数 seq_length = 32 vocab_size = 128 num_samples = 10000 # 生成模拟训练序列(真实音乐数据可替换为MIDI转换后的序列) X_train = np.random.randint(0, vocab_size, (num_samples, seq_length)) # 转成one-hot编码,适配模型输入 X_train_onehot = to_categorical(X_train, num_classes=vocab_size)
2. 构建生成器(Generator)
生成器输入随机噪声,输出符合时序逻辑的假序列:
def build_generator(latent_dim, seq_length, vocab_size): # 输入:随机噪声向量 input_noise = Input(shape=(latent_dim,)) # 将噪声映射为序列的初始特征 x = Dense(seq_length * 128)(input_noise) x = Reshape((seq_length, 128))(x) # 用LSTM捕捉时序依赖 x = LSTM(256, return_sequences=True)(x) x = LSTM(256, return_sequences=True)(x) # 输出每个时间步的音符概率分布 output = Dense(vocab_size, activation='softmax')(x) model = Model(input_noise, output) return model # 噪声向量维度 latent_dim = 100 generator = build_generator(latent_dim, seq_length, vocab_size) generator.summary()
3. 构建判别器(Discriminator)
判别器输入序列(真实/生成),输出该序列为真实数据的概率:
def build_discriminator(seq_length, vocab_size): # 输入:one-hot编码的序列 input_seq = Input(shape=(seq_length, vocab_size)) # LSTM提取时序特征 x = LSTM(256, return_sequences=True)(input_seq) x = LSTM(256)(x) # 输出真假判断(0=假,1=真) output = Dense(1, activation='sigmoid')(x) model = Model(input_seq, output) model.compile(optimizer=Adam(learning_rate=0.0002, beta_1=0.5), loss='binary_crossentropy', metrics=['accuracy']) return model discriminator = build_discriminator(seq_length, vocab_size) discriminator.summary()
4. 构建完整GAN模型
将生成器和判别器组合,训练生成器时冻结判别器的权重:
def build_gan(generator, discriminator): # 冻结判别器,只训练生成器 discriminator.trainable = False # GAN输入:随机噪声 gan_input = Input(shape=(latent_dim,)) # 生成器输出假序列 generated_seq = generator(gan_input) # 判别器判断生成序列的真假 gan_output = discriminator(generated_seq) model = Model(gan_input, gan_output) model.compile(optimizer=Adam(learning_rate=0.0002, beta_1=0.5), loss='binary_crossentropy') return model gan = build_gan(generator, discriminator) gan.summary()
5. 训练GAN
训练过程分为两步:先训练判别器区分真假序列,再训练生成器欺骗判别器:
def train_gan(gan, generator, discriminator, X_train_onehot, epochs=50, batch_size=64): batch_count = X_train_onehot.shape[0] // batch_size for epoch in range(epochs): for batch in range(batch_count): # ---------------------- # 1. 训练判别器 # ---------------------- # 取真实序列 real_seq = X_train_onehot[batch*batch_size : (batch+1)*batch_size] # 生成假序列 noise = np.random.normal(0, 1, (batch_size, latent_dim)) fake_seq = generator.predict(noise, verbose=0) # 准备判别器训练数据 X_disc = np.concatenate([real_seq, fake_seq]) y_disc = np.concatenate([np.ones((batch_size, 1)), np.zeros((batch_size, 1))]) # 加入标签噪声(提升稳定性) y_disc += 0.05 * np.random.random(y_disc.shape) # 训练判别器 d_loss, d_acc = discriminator.train_on_batch(X_disc, y_disc) # ---------------------- # 2. 训练生成器 # ---------------------- # 生成噪声,目标是让判别器认为生成的序列是真实的(y=1) noise = np.random.normal(0, 1, (batch_size, latent_dim)) y_gan = np.ones((batch_size, 1)) # 训练生成器 g_loss = gan.train_on_batch(noise, y_gan) # 打印每轮训练结果 print(f"Epoch {epoch+1}/{epochs} | Discriminator Loss: {d_loss:.4f}, Acc: {d_acc:.4f} | Generator Loss: {g_loss:.4f}") # 启动训练 train_gan(gan, generator, discriminator, X_train_onehot, epochs=50)
6. 生成音乐序列
训练完成后,用生成器生成新的序列,再转换回MIDI格式(可借助music21或midiutil库):
def generate_music_seq(generator, latent_dim, seq_length, vocab_size): noise = np.random.normal(0, 1, (1, latent_dim)) generated_onehot = generator.predict(noise, verbose=0)[0] # 从概率分布中采样得到整数序列(对应音符) generated_seq = np.argmax(generated_onehot, axis=1) return generated_seq # 生成序列 new_music_seq = generate_music_seq(generator, latent_dim, seq_length, vocab_size) print("生成的音乐序列(整数编码):", new_music_seq)
适配真实音乐数据的提示
- MIDI转序列:用
music21库读取MIDI文件,将音符、时长、和弦映射为整数编码,统一序列长度(截断或补零)。 - 序列特征:可加入更多维度(比如节拍、力度),将输入改为多特征序列。
- 网络调整:如果序列很长,可改用双向LSTM或Transformer层提升时序建模能力。
内容的提问来源于stack exchange,提问作者Stanislav Tataren
相关产品推荐
相关产品推荐

