Logo elodees  elodees

Une IA bien-veillante pour un monde meilleur













Seuls les caractères alphabétiques accentués ou non ainsi que l'espace sont acceptés

Logo IA




Réseaux profond antagoniste génératif Wasserstein





Pas encore de compte ?

Inscrivez-vous pour accéder à tous les contenus




Le réseau profond antagoniste génératif Wasserstein est un réseau antagoniste génératif qui remplace le discriminateur du réseau antagoniste génératif par un calcul de distance dit de Wasserstein.

Le réseau antagoniste génératif de Wasserstein est une variante du réseau antagoniste génératif qui améliore la stabilité de l'apprentissage, élimine l'effondrement de mode et fournit des courbes d'apprentissage significatives utiles pour le débogage et les recherches d'hyperparamètres.

Comparé au discriminateur GAN d'origine, le discriminateur Wasserstein GAN fournit un meilleur signal d'apprentissage au générateur.

Cela permet à la formation d'être plus stable lorsque le générateur apprend des distributions dans des espaces de très grande dimension.



Implémentation de Tensorflow de Wasserstein GAN (et version améliorée dans wgan_v2)



Testé sous Anaconda et Python 3.7

import os
import time
import argparse
import importlib
import tensorflow as tf
import imageio
 
from visualize import *
 
 
class WassersteinGAN(object):
    def __init__(self, g_net, d_net, x_sampler, z_sampler, data, model, scale=10.0):
        self.model = model
        self.data = data
        self.g_net = g_net
        self.d_net = d_net
        self.x_sampler = x_sampler
        self.z_sampler = z_sampler
        self.x_dim = self.d_net.x_dim
        self.z_dim = self.g_net.z_dim
        self.x = tf.placeholder(tf.float32, [None, self.x_dim], name='x')
        self.z = tf.placeholder(tf.float32, [None, self.z_dim], name='z')
 
        self.x_ = self.g_net(self.z)
 
        self.d = self.d_net(self.x, reuse=False)
        self.d_ = self.d_net(self.x_)
 
        self.g_loss = tf.reduce_mean(self.d_)
        self.d_loss = tf.reduce_mean(self.d) - tf.reduce_mean(self.d_)
 
        epsilon = tf.random_uniform([], 0.0, 1.0)
        x_hat = epsilon * self.x + (1 - epsilon) * self.x_
        d_hat = self.d_net(x_hat)
 
        ddx = tf.gradients(d_hat, x_hat)[0]
        print(ddx.get_shape().as_list())
        ddx = tf.sqrt(tf.reduce_sum(tf.square(ddx), axis=1))
        ddx = tf.reduce_mean(tf.square(ddx - 1.0) * scale)
 
        self.d_loss = self.d_loss + ddx
 
        self.d_adam, self.g_adam = None, None
        with tf.control_dependencies(tf.get_collection(tf.GraphKeys.UPDATE_OPS)):
            self.d_adam = tf.train.AdamOptimizer(learning_rate=1e-4, beta1=0.5, beta2=0.9)\
                .minimize(self.d_loss, var_list=self.d_net.vars)
            self.g_adam = tf.train.AdamOptimizer(learning_rate=1e-4, beta1=0.5, beta2=0.9)\
                .minimize(self.g_loss, var_list=self.g_net.vars)
 
        gpu_options = tf.GPUOptions(allow_growth=True)
        self.sess = tf.Session(config=tf.ConfigProto(gpu_options=gpu_options))
 
    def train(self, batch_size=64, num_batches=1000000):
        plt.ion()
        self.sess.run(tf.global_variables_initializer())
        start_time = time.time()
        for t in range(0, num_batches):
            d_iters = 5
            #if t % 500 == 0 or t < 25:
            #     d_iters = 100
 
            for _ in range(0, d_iters):
                bx = self.x_sampler(batch_size)
                bz = self.z_sampler(batch_size, self.z_dim)
                self.sess.run(self.d_adam, feed_dict={self.x: bx, self.z: bz})
 
            bz = self.z_sampler(batch_size, self.z_dim)
            self.sess.run(self.g_adam, feed_dict={self.z: bz, self.x: bx})
 
            if t % 100 == 0:
                bx = self.x_sampler(batch_size)
                bz = self.z_sampler(batch_size, self.z_dim)
 
                d_loss = self.sess.run(
                    self.d_loss, feed_dict={self.x: bx, self.z: bz}
                )
                g_loss = self.sess.run(
                    self.g_loss, feed_dict={self.z: bz}
                )
                print('Iter [%8d] Time [%5.4f] d_loss [%.4f] g_loss [%.4f]' %
                        (t, time.time() - start_time, d_loss, g_loss))
 
            if t % 100 == 0:
                bz = self.z_sampler(batch_size, self.z_dim)
                bx = self.sess.run(self.x_, feed_dict={self.z: bz})
                bx = xs.data2img(bx)
                #fig = plt.figure(self.data + '.' + self.model)
                #grid_show(fig, bx, xs.shape)
                bx = grid_transform(bx, xs.shape)
                imageio.imwrite('logs/{}/{}.png'.format(self.data, t/100), bx)
                #fig.savefig('logs/{}/{}.png'.format(self.data, t/100))
 
 
if __name__ == '__main__':
    parser = argparse.ArgumentParser('')
    parser.add_argument('--data', type=str, default='mnist')
    parser.add_argument('--model', type=str, default='dcgan')
    parser.add_argument('--gpus', type=str, default='0')
    args = parser.parse_args()
    os.environ['CUDA_VISIBLE_DEVICES'] = args.gpus
    data = importlib.import_module(args.data)
    model = importlib.import_module(args.data + '.' + args.model)
    xs = data.DataSampler()
    zs = data.NoiseSampler()
    d_net = model.Discriminator()
    g_net = model.Generator()
    wgan = WassersteinGAN(g_net, d_net, xs, zs, args.data, args.model)
    wgan.train()
 


Vous devez créer les répertoires logs/mnist

Remplacer from scipy.misc import imsave par import imageio
Remplacer imsave par imageio.imwrite



Lignes de commande :

python wgan_v2.py --data mnist --model mlp --gpus 0
python wgan_v2.py --data mnist --model dcgan --gpus 0



model mlp

Tensorflow playground

model dcgan

Tensorflow playground


wgan - GitHub





Réseau de neurones


Ingénierie des données


Apprentissage profond

Apprentissage automatique












Bienvenu, je m’appelle Eric Soupet et je suis l'administrateur du site elodees.com. elodees.com est un état de l'art de l'Intelligence Artificielle et se veut collaboratif, vous pouvez dès à présent proposer du contenu tels que des articles, des événements, des tutoriels, ... alors n'hésitez pas !

Crédit des images de la plate-forme : Pixabay - Pixabay License | Pexels - Pexels License