(config)
| 50 | |
| 51 | |
| 52 | def create_dataset(config): |
| 53 | if 'GMM' in config['dataset']: |
| 54 | if config['dataset'] == 'GMM_iso': |
| 55 | # isotropic |
| 56 | gmm_model = GMM(torch.tensor([0.33, 0.33, 0.34]), |
| 57 | torch.tensor([[-5, -5], [5, -5], [0, 5]]), |
| 58 | torch.tensor([[1, 1], [1, 1], [1, 1]])).cuda() |
| 59 | else: |
| 60 | # anisotropic |
| 61 | gmm_model = GMM(torch.tensor([0.33, 0.33, 0.34]), |
| 62 | torch.tensor([[-5, -5], [5, -5], [0, 5]]), |
| 63 | torch.tensor([[1.25, 0.5], [1.25, 0.5], [0.5, |
| 64 | 1.25]])).cuda() |
| 65 | |
| 66 | vis_density_GMM(gmm_model, config) |
| 67 | samples = gmm_model.sampling(config['num_samples']) |
| 68 | vis_2D_samples(samples.cpu().numpy(), config, tags='ground_truth') |
| 69 | train_set = GMMDataset(samples) |
| 70 | elif config['dataset'] == 'MNIST': |
| 71 | train_set = datasets.MNIST('./data', |
| 72 | train=True, |
| 73 | download=True, |
| 74 | transform=transforms.Compose([ |
| 75 | transforms.ToTensor(), |
| 76 | transforms.Normalize(config['img_mean'], |
| 77 | config['img_std']) |
| 78 | ])) |
| 79 | elif config['dataset'] == 'CelebA': |
| 80 | train_set = datasets.CelebA('./data', |
| 81 | split='train', |
| 82 | download=False, |
| 83 | transform=transforms.Compose([ |
| 84 | transforms.CenterCrop( |
| 85 | config['crop_size']), |
| 86 | transforms.Resize(config['height']), |
| 87 | transforms.ToTensor(), |
| 88 | transforms.Normalize(config['img_mean'], |
| 89 | config['img_std']) |
| 90 | ])) |
| 91 | elif config['dataset'] == 'CelebA2K': |
| 92 | train_set = datasets.CelebA('./data', |
| 93 | split='train', |
| 94 | download=False, |
| 95 | transform=transforms.Compose([ |
| 96 | transforms.CenterCrop( |
| 97 | config['crop_size']), |
| 98 | transforms.Resize(config['height']), |
| 99 | transforms.ToTensor(), |
| 100 | transforms.Normalize(config['img_mean'], |
| 101 | config['img_std']) |
| 102 | ])) |
| 103 | train_set = torch.utils.data.Subset(train_set, range(2000)) |
| 104 | elif config['dataset'] == 'FashionMNIST': |
| 105 | train_set = datasets.FashionMNIST('./data', |
| 106 | train=True, |
| 107 | download=True, |
| 108 | transform=transforms.Compose([ |
| 109 | transforms.ToTensor(), |
no test coverage detected