TrickmanOff

code2

Dec 19th, 2023
243
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
text 5.30 KB | None | 0 0
  1. """
  2. This file is expecting the 'pipeline' package name to be defined
  3. This file completely defines the experiment to run
  4. """
  5. import argparse
  6. import contextlib
  7. from typing import Tuple, Generator, Optional, Dict, List
  8.  
  9. import torch
  10. import torch.utils.data
  11. from torchvision import transforms
  12.  
  13. from lib import data
  14. from lib import logger
  15. from lib.discriminators import DCDiscriminator, FixedDCDiscriminator
  16. from lib.gan import GAN
  17. from lib.generators import DCGenerator, FixedDCGenerator
  18. from lib.metrics import *
  19. from lib.normalization import apply_normalization, SpectralNormalizer, ABCASNormalizer
  20. from lib.predicates import TrainPredicate, IgnoreFirstNEpochsPredicate, EachNthEpochPredicate
  21. from lib.storage import ExperimentsStorage
  22. from lib.train import Stepper, GANEpochTrainer, GanTrainer
  23. from lib.wandb_logger import WandbCM
  24.  
  25.  
  26. def init_storage() -> ExperimentsStorage:
  27. # === config variables ===
  28. experiments_dir = 'experiments'
  29. checkpoint_filename = './training_checkpoint'
  30. model_state_filename = './model_state'
  31. # ========================
  32. return ExperimentsStorage(experiments_dir=experiments_dir, checkpoint_filename=checkpoint_filename,
  33. model_state_filename=model_state_filename)
  34.  
  35.  
  36. experiments_storage = init_storage()
  37.  
  38.  
  39. def form_metric() -> Metric:
  40. return MetricsSequence(
  41. CriticValuesDistributionMetric(values_cnt=1000),
  42. GeneratedImagesMetric(5, 5),
  43. SSIMGenSimilarity(values_cnt=25),
  44. FIDMetric(values_cnt=2000),
  45. SSIMEvalMetric(values_cnt=2000),
  46. )
  47.  
  48. def form_per_batch_metrics():
  49. return [
  50. SSIMMetric(),
  51. ]
  52.  
  53.  
  54. def form_metric_predicate() -> Optional[TrainPredicate]:
  55. return EachNthEpochPredicate(5)
  56.  
  57.  
  58. def form_dataset(data_dir: str, train: bool = False) -> torch.utils.data.Dataset:
  59. if data_dir is None:
  60. dataset = data.get_cats_faces_dataset(train=train, val_ratio=0.13, load_all_in_memory=False)
  61. else:
  62. dataset = data.get_simple_images_dataset(data_dir, train=train, val_ratio=0.13, load_all_in_memory=False)
  63. return data.UnifiedDatasetWrapper(dataset)
  64.  
  65.  
  66. def init_logger(model_name: str = ''):
  67. project_name = 'GAN-pokemons'
  68. config = logger.get_default_config()
  69. @contextlib.contextmanager
  70. def logger_cm():
  71. try:
  72. with WandbCM(project_name=project_name, experiment_id=model_name, config=config) as wandb_logger:
  73. yield wandb_logger
  74. finally:
  75. pass
  76. return logger_cm
  77.  
  78.  
  79. def form_gan_trainer(data_dir: str, model_name: str,
  80. gan_model: Optional[GAN] = None, n_epochs: int = 100,
  81. enable_logging: bool = True) -> Generator[Tuple[int, GAN], None, GAN]:
  82. """
  83. :return: a generator that yields (epoch number, gan_model after this epoch)
  84. """
  85. logger_cm_fn = init_logger(model_name) if enable_logging else None
  86. metric = form_metric()
  87. metric_predicate = form_metric_predicate()
  88.  
  89. train_dataset = form_dataset(data_dir, train=True)
  90. val_dataset = form_dataset(data_dir, train=False)
  91.  
  92. # for local testing
  93. # val_size = int(0.1 * len(val_dataset))
  94. # val_dataset = torch.utils.data.Subset(val_dataset, np.arange(val_size))
  95. # -------
  96. noise_dimension = 100
  97.  
  98. def uniform_noise_generator(n: int, seed=None) -> torch.Tensor:
  99. gen = None if seed is None else torch.manual_seed(seed)
  100. return 2*torch.rand(size=(n, noise_dimension), generator=gen) - 1 # [-1, 1]
  101.  
  102. latent_channels = 512
  103. image_channels = 3
  104.  
  105. generator = FixedDCGenerator(noise_dim=noise_dimension, latent_channels=latent_channels, image_channels=image_channels)
  106. discriminator = FixedDCDiscriminator(latent_channels=latent_channels, image_channels=image_channels)
  107. # discriminator = apply_normalization(discriminator, SpectralNormalizer)
  108. # discriminator = apply_normalization(discriminator, ABCASNormalizer)
  109.  
  110. regularizer = None
  111.  
  112. if gan_model is None:
  113. gan_model = GAN(generator, discriminator, uniform_noise_generator)
  114.  
  115. generator_stepper = Stepper(
  116. optimizer=torch.optim.RMSprop(generator.parameters(), lr=5e-4)
  117. )
  118.  
  119. discriminator_stepper = Stepper(
  120. optimizer=torch.optim.RMSprop(discriminator.parameters(), lr=1e-4)
  121. )
  122.  
  123. per_batch_metrics = form_per_batch_metrics()
  124. epoch_trainer = GANEpochTrainer(n_critic=1, batch_size=128, use_wgan_loss=False, per_batch_metrics=per_batch_metrics)
  125.  
  126. model_dir = experiments_storage.get_model_dir(model_name)
  127. trainer = GanTrainer(model_dir=model_dir, use_saved_checkpoint=False, save_checkpoint_once_in_epoch=10)
  128. train_gan_generator = trainer.train(gan_model=gan_model,
  129. train_dataset=train_dataset, val_dataset=val_dataset,
  130. generator_stepper=generator_stepper,
  131. critic_stepper=discriminator_stepper,
  132. epoch_trainer=epoch_trainer,
  133. n_epochs=n_epochs,
  134. metric=metric, metric_predicate=metric_predicate,
  135. logger_cm_fn=logger_cm_fn,
  136. regularizer=regularizer)
  137. return train_gan_generator
Advertisement
Add Comment
Please, Sign In to add comment