| 1942 | |
| 1943 | |
| 1944 | def 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, |