ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

基于WGAN-GP的1D轴承振动数据样本生成方法

基于WGAN-GP的1D轴承振动数据样本生成方法 如何制作使用WGAN-GP-1D轴承振动数据样本生成方法以下文章代码仅供参考。文章目录如何制作使用WGAN-GP-1D轴承振动数据样本生成方法环境设置数据预处理preprocess_for_wgan.py训练过程train_gan.py测试过程generate_gan.py环境设置数据预处理preprocess_data.py训练过程 (train_gan.py)WGAN-GP版本测试过程 (generate_gan.py)WGAN-GP-1D轴承振动数据样本生成方法西储大学数据集为例可替换自己的数据。包含训练过程的代码train_gan和基于训练好的权重参数文件进行测试的代码generate_gan。WGAN-GP网络可更换为DCGAN、WGAN、LSGAN、SNGAN等。实现基于WGAN-GP的1D轴承振动数据样本生成方法我们可以使用Python和深度学习库PyTorch。以下是一个详细的示例代码包括训练过程train_gan.py和基于训练好的权重参数文件进行测试的代码generate_gan.py。将使用西储大学的数据集作为示例。、以下文章及代码仅供参考。环境设置确保你已经安装了以下依赖pipinstalltorch torchvision numpy scipy matplotlib数据预处理首先我们需要对数据进行预处理。假设你的数据集已经保存为.mat文件并且已经加载到了一个列表中。preprocess_for_wgan.pyimportscipy.ioassioimportnumpyasnpdefload_data(file_path): Load data from .mat file. datasio.loadmat(file_path)returndata[DE_time]defpreprocess_data(data,imbalance_ratio100): Preprocess the data for WGAN-GP. # Normalize the datadata(data-np.min(data))/(np.max(data)-np.min(data))# Create imbalanced datasetifimbalance_ratio1:normal_datadata[:int(len(data)/imbalance_ratio)]fault_datadata[int(len(data)/imbalance_ratio):]datanp.concatenate([normal_data,fault_data])returndata# Example usagefile_pathpath/to/your/data.matdataload_data(file_path)processed_datapreprocess_data(data,imbalance_ratio100)训练过程接下来是训练过程的代码。train_gan.pyimporttorchimporttorch.nnasnnimporttorch.optimasoptimfromtorch.utils.dataimportDataLoader,TensorDatasetimportnumpyasnpimportmatplotlib.pyplotaspltclassGenerator(nn.Module):def__init__(self,input_dim,output_dim):super(Generator,self).__init__()self.modelnn.Sequential(nn.Linear(input_dim,256),nn.ReLU(),nn.Linear(256,512),nn.ReLU(),nn.Linear(512,output_dim),nn.Tanh())defforward(self,x):returnself.model(x)classDiscriminator(nn.Module):def__init__(self,input_dim):super(Discriminator,self).__init__()self.modelnn.Sequential(nn.Linear(input_dim,512),nn.ReLU(),nn.Linear(512,256),nn.ReLU(),nn.Linear(256,1),nn.Sigmoid())defforward(self,x):returnself.model(x)defgradient_penalty(real_samples,fake_samples,discriminator):Calculates the gradient penalty.alphatorch.rand(real_samples.size(0),1)alphaalpha.expand_as(real_samples).cuda()interpolatesalpha*real_samples((1-alpha)*fake_samples)interpolatesinterpolates.cuda().requires_grad_(True)d_interpolatesdiscriminator(interpolates)gradientstorch.autograd.grad(outputsd_interpolates,inputsinterpolates,grad_outputstorch.ones(d_interpolates.size()).cuda(),create_graphTrue,retain_graphTrue,only_inputsTrue)[0]gradient_penalty((gradients.norm(2,dim1)-1)**2).mean()returngradient_penaltydeftrain_gan(data,epochs100,batch_size64,lr0.0001,lambda_gp10):devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)# Data preprocessingdatatorch.tensor(data,dtypetorch.float32).to(device)datasetTensorDataset(data)dataloaderDataLoader(dataset,batch_sizebatch_size,shuffleTrue)# Initialize generator and discriminatorinput_dim100# Latent space dimensionoutput_dimdata.shape[1]# Output dimension (same as data dimension)generatorGenerator(input_dim,output_dim).to(device)discriminatorDiscriminator(output_dim).to(device)# Optimizersoptimizer_Goptim.Adam(generator.parameters(),lrlr,betas(0.5,0.999))optimizer_Doptim.Adam(discriminator.parameters(),lrlr,betas(0.5,0.999))forepochinrange(epochs):fori,(real_samples,)inenumerate(dataloader):# Train Discriminatorreal_samplesreal_samples.to(device)ztorch.randn(batch_size,input_dim).to(device)fake_samplesgenerator(z)real_loss-torch.mean(discriminator(real_samples))fake_losstorch.mean(discriminator(fake_samples))gpgradient_penalty(real_samples,fake_samples,discriminator)d_lossreal_lossfake_losslambda_gp*gp optimizer_D.zero_grad()d_loss.backward(retain_graphTrue)optimizer_D.step()# Train Generatorztorch.randn(batch_size,input_dim).to(device)fake_samplesgenerator(z)g_loss-torch.mean(discriminator(fake_samples))optimizer_G.zero_grad()g_loss.backward()optimizer_G.step()print(fEpoch [{epoch1}/{epochs}], D Loss:{d_loss.item():.4f}, G Loss:{g_loss.item():.4f})# Save the trained modelstorch.save(generator.state_dict(),generator.pth)torch.save(discriminator.state_dict(),discriminator.pth)if__name____main__:file_pathpath/to/your/data.matdatapreprocess_data(load_data(file_path),imbalance_ratio100)train_gan(data,epochs100,batch_size64,lr0.0001,lambda_gp10)测试过程接下来是基于训练好的权重参数文件进行测试的代码。generate_gan.pyimporttorchimportnumpyasnpimportmatplotlib.pyplotaspltclassGenerator(nn.Module):def__init__(self,input_dim,output_dim):super(Generator,self).__init__()self.modelnn.Sequential(nn.Linear(input_dim,256),nn.ReLU(),nn.Linear(256,512),nn.ReLU(),nn.Linear(512,output_dim),nn.Tanh())defforward(self,x):returnself.model(x)defgenerate_gan(num_samples100,checkpointgenerator.pth):devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)# Load the trained generator modelinput_dim100# Latent space dimensionoutput_dim1024# Output dimension (same as data dimension)generatorGenerator(input_dim,output_dim).to(device)generator.load_state_dict(torch.load(checkpoint))generator.eval()# Generate samplesztorch.randn(num_samples,input_dim).to(device)generated_samplesgenerator(z).detach().cpu().numpy()# Plot the generated samplesplt.figure(figsize(10,6))plt.plot(generated_samples[0],labelGenerated Sample)plt.title(Generated Bearing Vibration Data)plt.xlabel(Time)plt.ylabel(Vibration Amplitude)plt.legend()plt.show()if__name____main__:generate_gan(num_samples100,checkpointgenerator.pth)该系统可以使用WGAN-GP或可选的DCGAN、WGAN、LSGAN、SNGAN等生成指定类型的故障轴承振动数据并能够根据不平衡比率调整生成的数据集以下是详细的示例代码。将分别提供训练过程(train_gan.py)和基于训练好的权重参数文件进行测试的代码(generate_gan.py)。环境设置确保你已经安装了必要的库pipinstalltorch torchvision numpy scipy matplotlib数据预处理首先我们需要对西储大学的轴承振动数据进行加载和预处理。这里假设你的数据是以.mat格式存储的。preprocess_data.pyimportscipy.ioassioimportnumpyasnpdefload_and_preprocess_data(file_path,fault_typeDE_time,imbalance_ratio100): Load and preprocess the bearing vibration data. :param file_path: Path to the .mat file containing the dataset. :param fault_type: The type of fault to extract from the dataset (e.g., DE_time). :param imbalance_ratio: Ratio for creating imbalanced dataset. :return: Preprocessed data. # Load datadatasio.loadmat(file_path)datadata[fault_type].flatten()# Normalize datadata(data-np.min(data))/(np.max(data)-np.min(data))# Create an imbalanced dataset if neededifimbalance_ratio1:split_pointlen(data)//imbalance_ratio normal_datadata[:split_point]fault_datadata[split_point:]datanp.concatenate([normal_data,fault_data])returndata训练过程 (train_gan.py)WGAN-GP版本importtorchfromtorchimportnnfromtorch.utils.dataimportDataLoader,TensorDatasetimportnumpyasnpfrompreprocess_dataimportload_and_preprocess_dataclassGenerator(nn.Module):def__init__(self,input_dim,output_dim):super(Generator,self).__init__()self.modelnn.Sequential(nn.Linear(input_dim,128),nn.ReLU(),nn.Linear(128,256),nn.ReLU(),nn.Linear(256,output_dim),nn.Tanh())defforward(self,x):returnself.model(x)classCritic(nn.Module):# In WGAN-GP, discriminator is referred to as criticdef__init__(self,input_dim):super(Critic,self).__init__()self.modelnn.Sequential(nn.Linear(input_dim,256),nn.ReLU(),nn.Linear(256,128),nn.ReLU(),nn.Linear(128,1))defforward(self,x):returnself.model(x)deftrain_wgan_gp(data,epochs100,batch_size64,lr1e-4,lambda_gp10):devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)# Prepare datadatatorch.tensor(data,dtypetorch.float32).to(device)datasetTensorDataset(data)dataloaderDataLoader(dataset,batch_sizebatch_size,shuffleTrue)# Initialize generator and criticinput_dim100# Latent space dimensionoutput_dimdata.shape[1]# Output dimension (same as data dimension)generatorGenerator(input_dim,output_dim).to(device)criticCritic(output_dim).to(device)# Optimizersoptimizer_Gtorch.optim.Adam(generator.parameters(),lrlr,betas(0.5,0.9))optimizer_Ctorch.optim.Adam(critic.parameters(),lrlr,betas(0.5,0.9))forepochinrange(epochs):fori,(real_samples,)inenumerate(dataloader):# Train Criticreal_samplesreal_samples.to(device)ztorch.randn(batch_size,input_dim,devicedevice)fake_samplesgenerator(z)gpcompute_gradient_penalty(critic,real_samples.data,fake_samples.data)critic_loss-(torch.mean(critic(real_samples))-torch.mean(critic(fake_samples)))lambda_gp*gp optimizer_C.zero_grad()critic_loss.backward(retain_graphTrue)optimizer_C.step()# Train Generator every n_critic stepsifi%50:fake_samplesgenerator(z)generator_loss-torch.mean(critic(fake_samples))optimizer_G.zero_grad()generator_loss.backward()optimizer_G.step()print(fEpoch [{epoch1}/{epochs}], C Loss:{critic_loss.item():.4f}, G Loss:{generator_loss.item():.4f})torch.save(generator.state_dict(),wgan_gp_generator.pth)defcompute_gradient_penalty(critic,real_samples,fake_samples):Calculates the gradient penalty.alphatorch.rand(real_samples.size(0),1).to(real_samples.device)interpolates(alpha*real_samples((1-alpha)*fake_samples)).requires_grad_(True)d_interpolatescritic(interpolates)faketorch.ones(d_interpolates.size()).to(real_samples.device)gradientstorch.autograd.grad(outputsd_interpolates,inputsinterpolates,grad_outputsfake,create_graphTrue,retain_graphTrue,only_inputsTrue)[0]gradient_penalty((gradients.norm(2,dim1)-1)**2).mean()returngradient_penaltyif__name____main__:file_pathpath/to/your/data.matdataload_and_preprocess_data(file_path,imbalance_ratio100)train_wgan_gp(data,epochs100,batch_size64,lr1e-4,lambda_gp10)测试过程 (generate_gan.py)importtorchimportmatplotlib.pyplotaspltfromGeneratorimportGenerator# Assuming you have saved the Generator class definition in a separate filedefgenerate_samples(checkpointwgan_gp_generator.pth,num_samples100,input_dim100):devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)# Load the trained generator modeloutput_dim1024# Example output dimension, should match your actual data dimensiongeneratorGenerator(input_dim,output_dim).to(device)generator.load_state_dict(torch.load(checkpoint))generator.eval()# Generate samplesztorch.randn(num_samples,input_dim,devicedevice)generated_samplesgenerator(z).detach().cpu().numpy()# Plot the first generated sampleplt.figure(figsize(10,6))plt.plot(generated_samples[0],labelGenerated Sample)plt.title(Generated Bearing Vibration Data)plt.xlabel(Time)plt.ylabel(Vibration Amplitude)plt.legend()plt.show()if__name____main__:generate_samples(num_samples100,checkpointwgan_gp_generator.pth)同学呀你可长点心吧上述代码中的网络结构例如层数、每层神经元的数量等可以根据实际需求进行调整。此外对于不同的GAN变体如DCGAN、WGAN、LSGAN、SNGAN同学你呀需要调整损失函数以及可能的网络架构以适应特定模型的要求。这些变化主要集中在Generator和Critic或Discriminator类的定义及其相应的训练过程中。
返回列表