import torch from giant.constants import COND_DIM, K_MAX, SEC_DIM, SEC_SLOT_DIM, X_DIM from giant.model.network import ( Critic, SecondaryCritic, WGANGenerator, WGANSecondaryGenerator, ) from giant.model.wgan import critic_loss, generator_loss, gradient_penalty from giant.sample import sample_secondaries_wgan, sample_wgan def _cond(B=8): cond_cont = torch.randn(B, COND_DIM) cond_cat = torch.zeros(B, 2, dtype=torch.long) return cond_cont, cond_cat def _small_generator(): return WGANGenerator( pdg_vocab=3, mat_vocab=2, hidden_dim=32, n_blocks=2, noise_dim=8 ) def _small_critic(): return Critic(pdg_vocab=3, mat_vocab=2, hidden_dim=32, n_blocks=2) def _small_sec_generator(): return WGANSecondaryGenerator( pdg_vocab=3, mat_vocab=2, hidden_dim=32, n_blocks=2, noise_dim=8 ) def _small_sec_critic(): return SecondaryCritic(pdg_vocab=3, mat_vocab=2, hidden_dim=32, n_blocks=2) def _mask(B, n_sec): sec_mask = torch.arange(K_MAX).unsqueeze(0) < n_sec.unsqueeze(1) return sec_mask.unsqueeze(-1).expand(-1, -1, SEC_SLOT_DIM).reshape(B, -1).float() # --- Stage-1 generator/critic --- def test_wgan_generator_output_shape(): B = 8 model = _small_generator() cond_cont, cond_cat = _cond(B) z = torch.randn(B, model.noise_dim) out = model(z, cond_cont, cond_cat) assert out.shape == (B, X_DIM) def test_wgan_generator_predict_n_sec_shape(): B = 6 model = _small_generator() cond_cont, cond_cat = _cond(B) logits = model.predict_n_sec(cond_cont, cond_cat) assert logits.shape == (B, K_MAX + 1) def test_wgan_generator_gradients_flow(): B = 4 model = _small_generator() cond_cont, cond_cat = _cond(B) z = torch.randn(B, model.noise_dim) gen_loss = model(z, cond_cont, cond_cat).sum() nsec_loss = model.predict_n_sec(cond_cont, cond_cat).sum() (gen_loss + nsec_loss).backward() for name, p in model.named_parameters(): assert p.grad is not None, f"no grad for {name}" def test_critic_output_shape(): B = 8 critic = _small_critic() cond_cont, cond_cat = _cond(B) x = torch.randn(B, X_DIM) out = critic(x, cond_cont, cond_cat) assert out.shape == (B,) def test_sample_wgan_shape(): B = 6 model = _small_generator() cond_cont, cond_cat = _cond(B) sample, n_sec = sample_wgan(model, cond_cont, cond_cat) assert sample.shape == (B, X_DIM) assert n_sec.shape == (B,) # --- Stage-2 generator/critic --- def test_wgan_secondary_generator_output_shape(): B = 8 model = _small_sec_generator() cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) z = torch.randn(B, model.noise_dim) out = model(z, cond_cont, cond_cat, stage1_out) assert out.shape == (B, SEC_DIM) def test_secondary_critic_output_shape(): B = 8 critic = _small_sec_critic() cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) x = torch.randn(B, SEC_DIM) out = critic(x, cond_cont, cond_cat, stage1_out) assert out.shape == (B,) def test_sample_secondaries_wgan_shape(): B = 5 model = _small_sec_generator() cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) n_sec_pred = torch.randint(0, K_MAX, (B,)) sec_cont, sec_phys, sec_valid = sample_secondaries_wgan( model, cond_cont, cond_cat, stage1_out, n_sec_pred ) assert sec_cont.shape == (B, K_MAX, 4) assert sec_phys.shape == (B, K_MAX, 2) assert sec_valid.shape == (B, K_MAX) # --- Losses --- def test_gradient_penalty_nonneg(): B = 8 critic = _small_critic() cond_cont, cond_cat = _cond(B) real = torch.randn(B, X_DIM) fake = torch.randn(B, X_DIM) gp = gradient_penalty(lambda x: critic(x, cond_cont, cond_cat), real, fake) assert gp.item() >= 0.0 assert gp.shape == () def test_gradient_penalty_masked(): B = 8 sec_critic = _small_sec_critic() cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) n_sec = torch.randint(0, K_MAX, (B,)) mask = _mask(B, n_sec) real = torch.randn(B, SEC_DIM) * mask fake = torch.randn(B, SEC_DIM) * mask gp = gradient_penalty( lambda x: sec_critic(x, cond_cont, cond_cat, stage1_out), real, fake, mask=mask ) assert gp.item() >= 0.0 def test_critic_loss_scalar_and_grad(): B = 8 critic = _small_critic() cond_cont, cond_cat = _cond(B) real = torch.randn(B, X_DIM) fake = torch.randn(B, X_DIM) loss = critic_loss( lambda x: critic(x, cond_cont, cond_cat), real, fake.detach(), gp_weight=10.0 ) assert loss.shape == () loss.backward() assert any(p.grad is not None for p in critic.parameters()) def test_generator_loss_scalar_and_grad(): B = 4 generator = _small_generator() critic = _small_critic() cond_cont, cond_cat = _cond(B) z = torch.randn(B, generator.noise_dim) fake = generator(z, cond_cont, cond_cat) loss = generator_loss(lambda x: critic(x, cond_cont, cond_cat), fake) assert loss.shape == () loss.backward() assert any(p.grad is not None for p in generator.parameters())