Antibody_deep_learning / scripts /07_train_gan.py
anzhi2710gmailcom's picture
Upload folder using huggingface_hub
fe8e241 verified
Raw
History Blame Contribute Delete
4.83 kB
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()