You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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)

适配真实音乐数据的提示

  1. MIDI转序列:用music21库读取MIDI文件,将音符、时长、和弦映射为整数编码,统一序列长度(截断或补零)。
  2. 序列特征:可加入更多维度(比如节拍、力度),将输入改为多特征序列。
  3. 网络调整:如果序列很长,可改用双向LSTM或Transformer层提升时序建模能力。

内容的提问来源于stack exchange,提问作者Stanislav Tataren

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.01 21:20:33