MCPcopy Create free account
hub / github.com/ddbourgin/numpy-ml / WGAN_GP_tf

Function WGAN_GP_tf

numpy_ml/tests/nn_torch_models.py:1944–2189  ·  view source on GitHub ↗
(X, lambda_, params, batch_size)

Source from the content-addressed store, hash-verified

1942
1943
1944def WGAN_GP_tf(X, lambda_, params, batch_size):
1945 tf.compat.v1.disable_eager_execution()
1946
1947 batch_size = X.shape[0]
1948
1949 # get alpha value
1950 n_steps = params["n_steps"]
1951 c_updates_per_epoch = params["c_updates_per_epoch"]
1952 alpha = tf.convert_to_tensor(params["alpha"], dtype="float32")
1953
1954 X_real = tf.compat.v1.placeholder(tf.float32, shape=[None, params["n_in"]])
1955 X_fake, G_out_X_fake, G_weights = Generator(batch_size, X_real, params)
1956
1957 Y_real, C_out_Y_real, C_Y_real_weights = Discriminator(X_real, params)
1958 Y_fake, C_out_Y_fake, C_Y_fake_weights = Discriminator(X_fake, params)
1959
1960 # WGAN loss
1961 mean_fake = tf.reduce_mean(Y_fake)
1962 mean_real = tf.reduce_mean(Y_real)
1963
1964 C_loss = tf.reduce_mean(Y_fake) - tf.reduce_mean(Y_real)
1965 G_loss = -tf.reduce_mean(Y_fake)
1966
1967 # WGAN gradient penalty
1968 X_interp = alpha * X_real + ((1 - alpha) * X_fake)
1969 Y_interp, C_out_Y_interp, C_Y_interp_weights = Discriminator(X_interp, params)
1970 gradInterp = tf.gradients(Y_interp, [X_interp])[0]
1971
1972 norm_gradInterp = tf.sqrt(
1973 tf.compat.v1.reduce_sum(tf.square(gradInterp), reduction_indices=[1])
1974 )
1975 gradient_penalty = tf.reduce_mean((norm_gradInterp - 1) ** 2)
1976 C_loss += lambda_ * gradient_penalty
1977
1978 # extract gradient of Y_interp wrt. each layer output in critic
1979 C_bwd_Y_interp = {}
1980 for k, v in C_out_Y_interp.items():
1981 C_bwd_Y_interp[k] = tf.gradients(Y_interp, [v])[0]
1982
1983 C_bwd_W = {}
1984 for k, v in C_Y_interp_weights.items():
1985 C_bwd_W[k] = tf.gradients(C_loss, [v])[0]
1986
1987 # get gradients
1988 dC_Y_fake = tf.gradients(C_loss, [Y_fake])[0]
1989 dC_Y_real = tf.gradients(C_loss, [Y_real])[0]
1990 dC_gradInterp = tf.gradients(C_loss, [gradInterp])[0]
1991 dG_Y_fake = tf.gradients(G_loss, [Y_fake])[0]
1992
1993 with tf.compat.v1.Session() as session:
1994 session.run(tf.compat.v1.global_variables_initializer())
1995
1996 for iteration in range(n_steps):
1997 # Train critic
1998 for i in range(c_updates_per_epoch):
1999 _data = X
2000 (
2001 _alpha,

Callers 1

test_WGAN_GPFunction · 0.85

Calls 3

GeneratorFunction · 0.85
DiscriminatorFunction · 0.85
gradientsMethod · 0.45

Tested by 1

test_WGAN_GPFunction · 0.68