import argparse import os import numpy as np import tensorflow as tf os.environ.setdefault("TF_FORCE_GPU_ALLOW_GROWTH", "true") def build_generator(latent_dim=100): inputs = tf.keras.Input(shape=(latent_dim,)) x = tf.keras.layers.Dense(16 * 11 * 128)(inputs) x = tf.keras.layers.LeakyReLU()(x) x = tf.keras.layers.Reshape((16, 11, 128))(x) x = tf.keras.layers.Conv2D(64, (2, 2), padding="same")(x) x = tf.keras.layers.LeakyReLU()(x) x = tf.keras.layers.Conv2DTranspose(32, (2, 2), strides=2, padding="same")(x) x = tf.keras.layers.LeakyReLU()(x) x = tf.keras.layers.Conv2D(128, (5, 6), padding="same")(x) x = tf.keras.layers.LeakyReLU()(x) outputs = tf.keras.layers.Conv2D(1, (6, 6), activation="tanh", padding="same")(x) return tf.keras.Model(inputs, outputs, name="generator") def build_discriminator(): inputs = tf.keras.Input(shape=(32, 22, 1)) x = tf.keras.layers.Conv2D(96, 3)(inputs) x = tf.keras.layers.LeakyReLU()(x) x = tf.keras.layers.Flatten()(x) x = tf.keras.layers.Dropout(0.3)(x) outputs = tf.keras.layers.Dense(1, activation="sigmoid")(x) return tf.keras.Model(inputs, outputs, name="discriminator") def train_gan(model_id, rounds=100, batch_size=20, latent_dim=100, seed=42): np.random.seed(seed + model_id) tf.random.set_seed(seed + model_id) data_file = f"model/GAN/seq_encoded_{model_id:02d}.npz" data = np.load(data_file, allow_pickle=True) real_data = data["x"].astype("float32") group_name = str(data["name"]) print(f"\n==== GAN model {model_id}: {group_name} ====") print("real_data:", real_data.shape) print("GPUs:", tf.config.list_physical_devices("GPU")) with tf.device("/GPU:0"): generator = build_generator(latent_dim) discriminator = build_discriminator() discriminator.compile( optimizer=tf.keras.optimizers.RMSprop(), loss="binary_crossentropy", ) discriminator.trainable = False gan_input = tf.keras.Input(shape=(latent_dim,)) gan_output = discriminator(generator(gan_input)) gan = tf.keras.Model(gan_input, gan_output, name="gan") gan.compile( optimizer=tf.keras.optimizers.RMSprop(), loss="binary_crossentropy", ) dloss = [] gloss = [] for step in range(1, rounds + 1): noise = np.random.normal(size=(batch_size, latent_dim)).astype("float32") fake = generator.predict(noise, verbose=0) idx = np.random.choice(real_data.shape[0], size=batch_size, replace=True) real = real_data[idx] both = np.concatenate([fake, real], axis=0).astype("float32") labels_fake = np.random.uniform(0.9, 1.0, size=(batch_size, 1)).astype("float32") labels_real = np.random.uniform(0.0, 0.1, size=(batch_size, 1)).astype("float32") labels = np.concatenate([labels_fake, labels_real], axis=0) discriminator.trainable = True d_loss = discriminator.train_on_batch(both, labels) noise = np.random.normal(size=(batch_size, latent_dim)).astype("float32") fake_as_real = np.random.uniform(0.0, 0.1, size=(batch_size, 1)).astype("float32") discriminator.trainable = False g_loss = gan.train_on_batch(noise, fake_as_real) dloss.append(float(d_loss)) gloss.append(float(g_loss)) if step == 1 or step % 10 == 0 or step == rounds: print(f"step {step:03d}/{rounds} dloss={dloss[-1]:.6f} gloss={gloss[-1]:.6f}") out_dir = f"weight/GAN/GAN_model_{model_id}_dcu" generator.save(out_dir) #generator.save(out_dir + ".keras") loss_file = f"weight/GAN/GAN_model_{model_id}_dcu_loss.npz" np.savez_compressed( loss_file, dloss=np.array(dloss, dtype="float32"), gloss=np.array(gloss, dtype="float32"), model_id=model_id, group_name=group_name, ) print("saved generator:", out_dir) print("saved loss:", loss_file) def main(): parser = argparse.ArgumentParser() parser.add_argument("--model-id", type=int, required=True, help="GAN model id, 1-15") parser.add_argument("--rounds", type=int, default=100) parser.add_argument("--batch-size", type=int, default=20) parser.add_argument("--latent-dim", type=int, default=100) parser.add_argument("--seed", type=int, default=42) args = parser.parse_args() if args.model_id < 1 or args.model_id > 15: raise ValueError("--model-id must be between 1 and 15") train_gan( model_id=args.model_id, rounds=args.rounds, batch_size=args.batch_size, latent_dim=args.latent_dim, seed=args.seed, ) if __name__ == "__main__": main()